diff --git a/posthog/dags/locations/signals.py b/posthog/dags/locations/signals.py index 07c386865bb9..82db9513636e 100644 --- a/posthog/dags/locations/signals.py +++ b/posthog/dags/locations/signals.py @@ -32,6 +32,7 @@ ], jobs=[ inbox_ranking_dataset.inbox_ranking_dataset_job, + inbox_ranking_dataset.inbox_ranking_labels_refresh_job, inbox_ranking_training.inbox_ranking_training_job, inbox_ranking_shadow.inbox_ranking_shadow_job, ], @@ -40,6 +41,7 @@ inbox_ranking_training.inbox_ranking_training_schedule, inbox_ranking_shadow.inbox_ranking_shadow_schedule, ], + sensors=[inbox_ranking_dataset.inbox_ranking_labels_refresh_sensor], loggers=loggers, resources=resources, ) diff --git a/posthog/settings/object_storage.py b/posthog/settings/object_storage.py index 9ab47f0b7836..643a03cb51cb 100644 --- a/posthog/settings/object_storage.py +++ b/posthog/settings/object_storage.py @@ -107,6 +107,9 @@ INBOX_RANKING_TRAINING_HOLDOUT_DAYS = get_from_env("INBOX_RANKING_TRAINING_HOLDOUT_DAYS", 7, type_cast=int) INBOX_RANKING_AUTO_PROMOTE = get_from_env("INBOX_RANKING_AUTO_PROMOTE", False, type_cast=str_to_bool) INBOX_RANKING_PROMOTION_MIN_DAYS = get_from_env("INBOX_RANKING_PROMOTION_MIN_DAYS", 3, type_cast=int) +# Labels refresh sensor (products/signals/dags/inbox_ranking/dataset): how many stale labels +# partitions one hourly tick rewrites after a FEATURE_SCHEMA_VERSION bump, newest first. +INBOX_RANKING_LABELS_REFRESH_MAX_RUNS = get_from_env("INBOX_RANKING_LABELS_REFRESH_MAX_RUNS", 6, type_cast=int) # The family whose champion the serving manifest serves. The scoring sweep reads the manifest # from the deployment's own object store, so this is the only place the served family is chosen. INBOX_RANKING_SERVED_FAMILY = os.getenv("INBOX_RANKING_SERVED_FAMILY", "report_embeddings") diff --git a/products/signals/dags/inbox_ranking/README.md b/products/signals/dags/inbox_ranking/README.md index 642dc6d68dc6..c38042977879 100644 --- a/products/signals/dags/inbox_ranking/README.md +++ b/products/signals/dags/inbox_ranking/README.md @@ -262,6 +262,8 @@ All reads route to the offline cluster replicas on Cloud (`etl_workload()`), car ## Operating it - Backfill any day range from the Dagster UI; partitions start 2026-04-01 (the label epoch). Every asset sits in the `inbox_ranking_etl` pool so concurrent partitions don't each start their own fleet-wide embeddings scan — the pool's limit is a Dagster deployment setting, provisioned with the bucket. +- A label column change needs a `FEATURE_SCHEMA_VERSION` bump, and no manual backfill. The labels asset stamps the version in each object's `feature-schema-version` metadata. Every hour, `inbox_ranking_labels_refresh_sensor` finds labels partitions in the training lookback whose stamp is missing or older, and runs `inbox_ranking_labels_refresh_job` on at most `INBOX_RANKING_LABELS_REFRESH_MAX_RUNS` (default 6) of them, newest first. A 60-day lookback is current again in about 10 hours. The sensor skips the newest day, which the daily schedule writes, and partitions with no labels object. It requests a partition at most once per schema version, so a failed refresh alerts and needs a person. The refresh covers labels only: report state reads current Postgres and embeddings have a TTL, so a rewrite of those is not point-in-time. `pairs_skipped_missing_label_columns` on `inbox_ranking_examples_built` counts the snapshot pairs each head lost to a missing label column while partitions are stale. +- Do not materialize these assets in-process in the code-location pod. Its memory limit (2Gi) is too small; launch a run instead. - Failures alert `#alerts-self-driving` (owner `team-self-driving`); assets retry twice with a 60s delay before failing a run. A UI-launched materialization runs under Dagster's implicit `__ASSET_JOB`, which carries no owner tag, so alert routing falls back to matching the `inbox_report_`, `inbox_signal_`, and `inbox_ranking_` asset-name prefixes. - Runtime budgets are per job (`dagster/max_runtime`): 3h for the dataset and training jobs, 1h for the shadow job. The 3h figure is what the dataset needs — its seven label streams run sequentially, each allowed up to 600s, and the join and S3 writes come after them. The shadow read is one day of two event families plus the scores objects in its lookback, so it gets an hour. diff --git a/products/signals/dags/inbox_ranking/common.py b/products/signals/dags/inbox_ranking/common.py index 706edee3baaa..31386b11fe0a 100644 --- a/products/signals/dags/inbox_ranking/common.py +++ b/products/signals/dags/inbox_ranking/common.py @@ -134,9 +134,19 @@ def serving_mirror_storage() -> ObjectStorage: SNAPSHOT_DATE_METADATA_KEY = "snapshot-date" ROW_COUNT_METADATA_KEY = "row-count" - - -def write_parquet(client, bucket: str, key: str, table: pa.Table, snapshot_date: str | None = None) -> None: +# The FEATURE_SCHEMA_VERSION the labels asset wrote an object under. The refresh sensor rewrites a +# partition whose stamp is missing or older, so a new label column reaches the whole lookback. +SCHEMA_VERSION_METADATA_KEY = "feature-schema-version" + + +def write_parquet( + client, + bucket: str, + key: str, + table: pa.Table, + snapshot_date: str | None = None, + schema_version: int | None = None, +) -> None: """Write one Parquet object at a deterministic key. Spooled to a temp file and uploaded with `upload_fileobj` rather than held as bytes for @@ -147,6 +157,8 @@ def write_parquet(client, bucket: str, key: str, table: pa.Table, snapshot_date: metadata = {ROW_COUNT_METADATA_KEY: str(table.num_rows)} if snapshot_date: metadata[SNAPSHOT_DATE_METADATA_KEY] = snapshot_date + if schema_version is not None: + metadata[SCHEMA_VERSION_METADATA_KEY] = str(schema_version) with tempfile.TemporaryFile() as spool: pq.write_table(table, spool, compression="zstd") spool.seek(0) @@ -178,6 +190,19 @@ def object_row_count(client, bucket: str, key: str) -> int | None: return int(stamped) if stamped is not None else None +def object_schema_version(client, bucket: str, key: str) -> int | None: + """The schema version stamped on an object at write time, or None when the object is missing + or predates the stamp.""" + try: + head = client.head_object(Bucket=bucket, Key=key) + except ClientError as error: + if error.response.get("Error", {}).get("Code") in ("404", "NoSuchKey", "NotFound"): + return None + raise + stamped = head.get("Metadata", {}).get(SCHEMA_VERSION_METADATA_KEY) + return int(stamped) if stamped is not None else None + + def partition_write_allowed(existing_row_count: int | None, row_count: int) -> bool: """Whether a re-run may overwrite a partition it has already written. diff --git a/products/signals/dags/inbox_ranking/dataset/dag.py b/products/signals/dags/inbox_ranking/dataset/dag.py index 1e57cca18453..cc0bb0ac5f4f 100644 --- a/products/signals/dags/inbox_ranking/dataset/dag.py +++ b/products/signals/dags/inbox_ranking/dataset/dag.py @@ -48,7 +48,7 @@ import json import datetime -from collections.abc import Iterator +from collections.abc import Collection, Iterator, Mapping from typing import Any, cast from django.db.models import Count, Min @@ -77,6 +77,7 @@ from products.signals.dags.inbox_ranking.common import ( DATASET_VERSION, HUMAN_ACTOR_KINDS, + PARQUET_PART_NAME, S3_BUCKET_ENV, dataset_bucket, dataset_unconfigured, @@ -85,6 +86,7 @@ latest_object_key, merge_emission_rows, object_row_count, + object_schema_version, object_snapshot_date, owner_tags, partition_def, @@ -867,7 +869,13 @@ def inbox_report_labels(context: dagster.AssetExecutionContext) -> None: rows = merge_label_streams(stream_rows, datetime.date.fromisoformat(partition_key)) bucket = dataset_bucket() key = partition_object_key(settings.INBOX_RANKING_DATASET_S3_PREFIX, LABELS_TABLE, partition_key) - write_parquet(s3_client(), bucket, key, pa.Table.from_pylist(rows, schema=LABELS_SCHEMA)) + write_parquet( + s3_client(), + bucket, + key, + pa.Table.from_pylist(rows, schema=LABELS_SCHEMA), + schema_version=FEATURE_SCHEMA_VERSION, + ) context.add_output_metadata( { "rows": dagster.MetadataValue.int(len(rows)), @@ -1097,3 +1105,95 @@ def inbox_ranking_dataset_schedule( # builds the partition it was scheduled for; run_key dedupes a re-evaluated tick. previous_day = context.scheduled_execution_time.date() - datetime.timedelta(days=1) return dagster.RunRequest(partition_key=previous_day.isoformat(), run_key=previous_day.isoformat()) + + +inbox_ranking_labels_refresh_job = dagster.define_asset_job( + name="inbox_ranking_labels_refresh_job", + selection=[LABELS_TABLE], + tags={**owner_tags, "dagster/max_runtime": str(3 * 60 * 60)}, +) + + +def stale_label_partitions( + stamps: Mapping[str, int | None], current: int, limit: int, requested: Collection[str] = () +) -> list[str]: + """The partitions to rewrite, newest first. `stamps` holds only partitions whose labels object + exists. A missing object has no state snapshot either, so a rewrite cannot make it an example. + `requested` are partitions already requested at `current`, in flight or failed.""" + stale = [ + partition + for partition, version in stamps.items() + if (version is None or version < current) and partition not in requested + ] + return sorted(stale, reverse=True)[:limit] + + +def label_refresh_window(today: datetime.date) -> list[str]: + """The training lookback, without the newest day. The daily schedule writes `today - 1`, so the + sensor never writes the same object at the same time.""" + start = max( + today - datetime.timedelta(days=1 + settings.INBOX_RANKING_TRAINING_LOOKBACK_DAYS), + partition_def.start.date(), + ) + end = today - datetime.timedelta(days=2) + return [(start + datetime.timedelta(days=offset)).isoformat() for offset in range((end - start).days + 1)] + + +def _existing_label_partitions(client, bucket: str, prefix: str) -> set[str]: + table_prefix = f"{prefix}/{LABELS_TABLE}/{DATASET_VERSION}/dt=" + partitions: set[str] = set() + for page in client.get_paginator("list_objects_v2").paginate(Bucket=bucket, Prefix=table_prefix): + for item in page.get("Contents", []): + partition, _, name = item["Key"].removeprefix(table_prefix).partition("/") + if name == PARQUET_PART_NAME: + partitions.add(partition) + return partitions + + +# A label column change bumps FEATURE_SCHEMA_VERSION, and the training lookback then holds +# partitions without that column. The heads that read it train on almost no examples until those +# partitions are rewritten. A partition is requested at most once per version: the cursor skips a +# run that is in flight or failed, so the next tick moves on to older partitions. A failed refresh +# alerts like any other run failure and needs a person. +@dagster.sensor( + job=inbox_ranking_labels_refresh_job, + minimum_interval_seconds=60 * 60, + default_status=dagster.DefaultSensorStatus.RUNNING + if settings.CLOUD_DEPLOYMENT == "US" + else dagster.DefaultSensorStatus.STOPPED, +) +def inbox_ranking_labels_refresh_sensor( + context: dagster.SensorEvaluationContext, +) -> dagster.SensorResult | dagster.SkipReason: + if dataset_unconfigured(): + return dagster.SkipReason(f"{S3_BUCKET_ENV} is not set; skipping until the dedicated bucket is provisioned") + cursor = json.loads(context.cursor) if context.cursor else {} + requested = set(cursor.get("requested", [])) if cursor.get("version") == FEATURE_SCHEMA_VERSION else set() + + client, bucket, prefix = s3_client(), dataset_bucket(), settings.INBOX_RANKING_DATASET_S3_PREFIX + existing = _existing_label_partitions(client, bucket, prefix) + stamps = { + partition: object_schema_version(client, bucket, partition_object_key(prefix, LABELS_TABLE, partition)) + for partition in label_refresh_window(datetime.datetime.now(datetime.UTC).date()) + if partition in existing + } + stale = stale_label_partitions(stamps, FEATURE_SCHEMA_VERSION, len(stamps)) + batch = stale_label_partitions( + stamps, FEATURE_SCHEMA_VERSION, settings.INBOX_RANKING_LABELS_REFRESH_MAX_RUNS, requested + ) + context.log.info(f"stale_label_partitions={len(stale)} requesting={len(batch)} (schema v{FEATURE_SCHEMA_VERSION})") + if not batch: + return dagster.SkipReason(f"stale_label_partitions={len(stale)}, none left to request") + return dagster.SensorResult( + run_requests=[ + dagster.RunRequest( + partition_key=partition, + run_key=f"{partition}-labels-v{FEATURE_SCHEMA_VERSION}", + tags=owner_tags, + ) + for partition in batch + ], + cursor=json.dumps( + {"version": FEATURE_SCHEMA_VERSION, "requested": sorted(requested.union(batch) & set(stale))} + ), + ) diff --git a/products/signals/dags/inbox_ranking/tests/test_dataset.py b/products/signals/dags/inbox_ranking/tests/test_dataset.py index 54e1fbd0e6cc..1bedea0f6942 100644 --- a/products/signals/dags/inbox_ranking/tests/test_dataset.py +++ b/products/signals/dags/inbox_ranking/tests/test_dataset.py @@ -107,6 +107,37 @@ def test_latest_advances_monotonically_and_backfills_never_clobber_it(existing, assert common.latest_is_stale(existing, partition_key) is expected +@pytest.mark.parametrize( + "stamps,requested,limit,expected", + [ + ({"2026-07-01": None, "2026-07-02": 8, "2026-07-03": 9}, (), 6, ["2026-07-02", "2026-07-01"]), + ({"2026-07-01": 10}, (), 6, []), + ({f"2026-07-{day:02d}": None for day in range(1, 11)}, (), 3, ["2026-07-10", "2026-07-09", "2026-07-08"]), + ({"2026-07-01": 8, "2026-07-02": 8, "2026-07-03": 8}, ("2026-07-03",), 1, ["2026-07-02"]), + ], +) +def test_stale_label_partitions_newest_first_capped_and_skips_requested(stamps, requested, limit, expected): + assert dag.stale_label_partitions(stamps, 9, limit, requested) == expected + + +class _MetadataS3: + def __init__(self) -> None: + self.metadata: dict[str, dict[str, str]] = {} + + def upload_fileobj(self, fileobj, bucket, key, ExtraArgs): + self.metadata[key] = ExtraArgs["Metadata"] + + def head_object(self, Bucket, Key): + return {"Metadata": self.metadata[Key]} + + +@pytest.mark.parametrize("schema_version", [9, None]) +def test_schema_version_stamp_round_trips(schema_version): + client = _MetadataS3() + common.write_parquet(client, "b", "k", pa.table({"x": [1]}), schema_version=schema_version) + assert common.object_schema_version(client, "b", "k") == schema_version + + @pytest.mark.parametrize( "existing,row_count,expected", [ diff --git a/products/signals/dags/inbox_ranking/tests/test_training.py b/products/signals/dags/inbox_ranking/tests/test_training.py index 92219f16842c..5a7ddbdf9447 100644 --- a/products/signals/dags/inbox_ranking/tests/test_training.py +++ b/products/signals/dags/inbox_ranking/tests/test_training.py @@ -1500,6 +1500,7 @@ def test_training_events_carry_the_dashboard_contract(monkeypatch): birth_day_positives=1, example_window_start=datetime.date(2026, 7, 1), example_cap_bound=True, + pairs_skipped_missing_label_columns=0, ) }, ), diff --git a/products/signals/dags/inbox_ranking/training/dag.py b/products/signals/dags/inbox_ranking/training/dag.py index b8271ad3ea0e..b8b5a559d553 100644 --- a/products/signals/dags/inbox_ranking/training/dag.py +++ b/products/signals/dags/inbox_ranking/training/dag.py @@ -507,6 +507,7 @@ def _write_examples( birth_day_positives=birth_day_positives(head_examples.examples), example_window_start=head_examples.window_start, example_cap_bound=head_examples.cap_bound, + pairs_skipped_missing_label_columns=head_examples.pairs_skipped_missing_label_columns, ) for name, head_examples in built.items() } @@ -524,6 +525,12 @@ def _write_examples( f"{feature_set.name}_{name}_birth_day_positives": dagster.MetadataValue.int(head_counts.birth_day_positives) for name, head_counts in counts.items() }, + **{ + f"{feature_set.name}_{name}_pairs_skipped_missing_label_columns": dagster.MetadataValue.int( + head_counts.pairs_skipped_missing_label_columns + ) + for name, head_counts in counts.items() + }, f"{feature_set.name}_s3_key": dagster.MetadataValue.text(f"s3://{bucket}/{key}"), } capture_training_events( diff --git a/products/signals/dags/inbox_ranking/training/examples.py b/products/signals/dags/inbox_ranking/training/examples.py index e0a7c6666ded..7391c19d2204 100644 --- a/products/signals/dags/inbox_ranking/training/examples.py +++ b/products/signals/dags/inbox_ranking/training/examples.py @@ -225,6 +225,9 @@ class HeadExamples: window_start: datetime.date | None # True when the row budget dropped at least one older day. cap_bound: bool + # Snapshot pairs dropped because one side lacks a label column the head reads. A high count + # with few positives means the labels partitions predate the current schema. + pairs_skipped_missing_label_columns: int def window(self) -> dict[str, object]: return { @@ -247,6 +250,22 @@ def build_head_examples( examples=_with_features(kept, snapshots, feature_set, extras), window_start=days.min().date() if len(days) else None, cap_bound=len(kept) < len(moments), + pairs_skipped_missing_label_columns=pairs_missing_label_columns(snapshots, head), + ) + + +def _label_columns_readable(now: Snapshot, later: Snapshot, head: Head) -> bool: + return all(column in now.labels and column in later.labels for column in head.label_columns) + + +def pairs_missing_label_columns(snapshots: Mapping[datetime.date, Snapshot], head: Head) -> int: + """The (snapshot, `horizon_days`-later snapshot) pairs `example_moments` skips for a missing + label column.""" + return sum( + 1 + for date, now in snapshots.items() + if (later := snapshots.get(date + datetime.timedelta(days=head.horizon_days))) is not None + and not _label_columns_readable(now, later, head) ) @@ -271,7 +290,7 @@ def example_moments( # only in the later snapshot (a column that entered the schema mid-window) would pass the # "not yet observed at now" guard below and mint an outcome from before `now` as a future # positive. Skip the pair when the head's label cannot be read from both snapshots. - if any(column not in now.labels or column not in later.labels for column in head.label_columns): + if not _label_columns_readable(now, later, head): continue ids = now.state.index.intersection(now.labels.index).intersection(later.labels.index) if len(ids) == 0: diff --git a/products/signals/dags/inbox_ranking/training/telemetry.py b/products/signals/dags/inbox_ranking/training/telemetry.py index dd39faa28460..2001b828ca19 100644 --- a/products/signals/dags/inbox_ranking/training/telemetry.py +++ b/products/signals/dags/inbox_ranking/training/telemetry.py @@ -72,6 +72,7 @@ class HeadExampleCounts: # so a chart shows when the budget starts to cut history. example_window_start: datetime.date | None example_cap_bound: bool + pairs_skipped_missing_label_columns: int def candidate_events(metadata: Mapping[str, Any]) -> list[TrainingEvent]: @@ -136,6 +137,7 @@ def examples_events( if counts.example_window_start else None, "example_cap_bound": counts.example_cap_bound, + "pairs_skipped_missing_label_columns": counts.pairs_skipped_missing_label_columns, }, ) for head, counts in per_head.items()