kaxil commented on code in PR #74407:
URL: https://github.com/apache/airflow/pull/74407#discussion_r4211011947


##########
RELEASE_NOTES.rst:
##########
@@ -24,6 +24,21 @@
 
 .. towncrier release notes start
 
+Airflow 3.4.0 (unreleased)
+--------------------------
+
+.. note::
+  Downgrading from this release to 3.4.0 can be slow, depending on the number 
of

Review Comment:
   "from this release to 3.4.0" should say 3.3.x, I think.
   
   The cost also isn't only proportional to rows written after the upgrade. On 
Postgres `_redirect_legacy` re-adds the `xcom` and 
`rendered_task_instance_fields` FKs as validated constraints, which scans the 
whole legacy table, and that can't be `NOT VALID` because a re-upgrade's 
`_check_source` requires a validated FK.
   
   The PR description says the caveats are documented here, but the note 
doesn't mention that XCom and rendered fields of every attempt except the 
latest are dropped, or that the downgrade can still refuse (orphaned archived 
attempts, owners whose coordinates moved, offline `--sql`). Could those go in 
too, so nobody finds out mid-rollback?



##########
airflow-core/src/airflow/migrations/versions/0142_3_4_0_unify_task_attempt_ownership.py:
##########
@@ -370,24 +370,244 @@ def upgrade():
         )
 
 
+def _build_task_instance_table():
+    return sa.table(
+        "task_instance",
+        sa.column("id"),
+        *(sa.column(c) for c in _COPY_COLUMNS),
+        sa.column("working_set", sa.Boolean()),
+        sa.column("dag_version_id"),
+    )
+
+
+def _build_legacy_owner_table():
+    return sa.table(
+        "legacy_task_data_owner", *(sa.column(c) for c in _COORDINATES), 
sa.column("task_instance_id")
+    )
+
+
+def _select_latest_attempts(ti):
+    newer = ti.alias("newer")
+    # An anti-join, because grouping by max(try_number) aggregates the whole 
table however few rows need copying.
+    return (
+        sa.select(ti.c.id, *(ti.c[c] for c in _COORDINATES))
+        .where(
+            ~sa.exists().where(
+                *(newer.c[c] == ti.c[c] for c in _COORDINATES), 
newer.c.try_number > ti.c.try_number
+            )
+        )
+        .subquery()
+    )
+
+
+def _has_duplicate_attempts(bind, ti) -> bool:
+    coordinates = [ti.c[c] for c in _COORDINATES]
+    for where, group_by in (
+        (ti.c.working_set.is_not(None), coordinates),
+        (sa.true(), [*coordinates, ti.c.try_number]),
+    ):
+        duplicated = 
sa.select(1).select_from(ti).where(where).group_by(*group_by).having(sa.func.count()
 > 1)
+        if bind.scalar(duplicated.limit(1)):
+            return True
+    return False
+
+
+def _select_archived_ids(ti):
+    return sa.select(ti.c.id).where(ti.c.working_set.is_(None))
+
+
+def _count_orphaned_archived_attempts(bind, ti) -> int:
+    current = ti.alias("current")
+    orphaned = (
+        sa.select(sa.func.count())
+        .select_from(ti)
+        .where(
+            ti.c.working_set.is_(None),
+            ~sa.exists().where(
+                *(current.c[c] == ti.c[c] for c in _COORDINATES), 
current.c.working_set.is_not(None)
+            ),
+        )
+    )
+    return bind.scalar(orphaned)
+
+
+def _has_moved_owners_with_legacy_rows(bind, ti) -> bool:
+    owner = _build_legacy_owner_table()
+    moved = sa.or_(*(owner.c[c] != ti.c[c] for c in _COORDINATES))
+    for name in ("xcom_v1", "rtif_v1"):
+        legacy = sa.table(name, *(sa.column(c) for c in _COORDINATES))
+        found = (
+            sa.select(1)
+            .select_from(owner.join(ti, ti.c.id == owner.c.task_instance_id))
+            .where(moved, sa.exists().where(*(legacy.c[c] == owner.c[c] for c 
in _COORDINATES)))
+        )
+        if bind.scalar(found.limit(1)):
+            return True
+    return False
+
+
+def _delete_moved_owners(ti):
+    owner = _build_legacy_owner_table()
+    op.execute(
+        owner.delete().where(
+            sa.exists().where(
+                ti.c.id == owner.c.task_instance_id,
+                sa.or_(*(ti.c[c] != owner.c[c] for c in _COORDINATES)),
+            )
+        )
+    )
+
+
+def _delete_legacy_rows_of_archived_attempts(ti):
+    owner = _build_legacy_owner_table()
+    archived_ids = _select_archived_ids(ti)
+    # Legacy rows owned by an archived attempt are hidden from its successor, 
so keeping them would expose them again.
+    for name in ("xcom_v1", "rtif_v1"):
+        legacy = sa.table(name, *(sa.column(c) for c in _COORDINATES))
+        op.execute(
+            legacy.delete().where(
+                sa.exists().where(
+                    *(owner.c[c] == legacy.c[c] for c in _COORDINATES),
+                    owner.c.task_instance_id.in_(archived_ids),
+                )
+            )
+        )
+
+
+def _move_owners_to_current_attempts(ti):
+    owner = _build_legacy_owner_table()
+    current = ti.alias("current")
+    archived_ids = _select_archived_ids(ti)
+    xcom_v2 = sa.table("xcom_v2", sa.column("task_instance_id"))
+    rtif_v2 = sa.table("rtif_v2", sa.column("task_instance_id"))
+    # Deleting archived attempts must not cascade into legacy rows through an 
owner that points at one.
+    op.execute(
+        owner.update()
+        .where(owner.c.task_instance_id.in_(archived_ids))
+        .values(
+            task_instance_id=sa.select(current.c.id)
+            .where(*(current.c[c] == owner.c[c] for c in _COORDINATES), 
current.c.working_set.is_not(None))
+            .scalar_subquery()
+        )
+    )
+    # Legacy tables reference owners by coordinates, so attempts created after 
the upgrade need an owner row first.
+    op.execute(
+        owner.insert().from_select(
+            [*_COORDINATES, "task_instance_id"],
+            sa.select(*(ti.c[c] for c in _COORDINATES), ti.c.id).where(
+                ti.c.working_set.is_not(None),
+                ti.c.id.in_(sa.select(xcom_v2.c.task_instance_id))
+                | ti.c.id.in_(sa.select(rtif_v2.c.task_instance_id)),
+                ~sa.exists().where(*(owner.c[c] == ti.c[c] for c in 
_COORDINATES)),
+            ),
+        )
+    )
+
+
+def _delete_replaced_legacy_rows(v1, v2, latest, *matching):
+    op.execute(
+        v1.delete().where(
+            sa.exists().where(
+                v2.c.task_instance_id == latest.c.id,
+                *(latest.c[c] == v1.c[c] for c in _COORDINATES),
+                *matching,
+            )
+        )
+    )
+
+
+def _copy_xcom_to_legacy(latest):
+    columns = ("key", "value", "timestamp", "dag_result", "mapped_length")
+    v2 = sa.table("xcom_v2", sa.column("task_instance_id"), *(sa.column(c) for 
c in columns))
+    v1 = sa.table("xcom_v1", *(sa.column(c) for c in (*_COORDINATES, *columns, 
"dag_run_id")))
+    dag_run = sa.table("dag_run", sa.column("id"), sa.column("dag_id"), 
sa.column("run_id"))
+    _delete_replaced_legacy_rows(v1, v2, latest, v2.c.key == v1.c.key)
+    source = v2.join(latest, v2.c.task_instance_id == latest.c.id).join(
+        dag_run, sa.and_(dag_run.c.dag_id == latest.c.dag_id, dag_run.c.run_id 
== latest.c.run_id)
+    )
+    op.execute(
+        v1.insert().from_select(
+            [*_COORDINATES, *columns, "dag_run_id"],
+            sa.select(
+                *(latest.c[c] for c in _COORDINATES), *(v2.c[c] for c in 
columns), dag_run.c.id
+            ).select_from(source),
+        )
+    )
+
+
+def _copy_rendered_fields_to_legacy(latest):
+    columns = ("rendered_fields", "k8s_pod_yaml")
+    v2 = sa.table("rtif_v2", sa.column("task_instance_id"), *(sa.column(c) for 
c in columns))
+    v1 = sa.table("rtif_v1", *(sa.column(c) for c in (*_COORDINATES, 
*columns)))
+    _delete_replaced_legacy_rows(v1, v2, latest)
+    op.execute(
+        v1.insert().from_select(
+            [*_COORDINATES, *columns],
+            sa.select(*(latest.c[c] for c in _COORDINATES), *(v2.c[c] for c in 
columns)).select_from(
+                v2.join(latest, v2.c.task_instance_id == latest.c.id)
+            ),
+        )
+    )
+
+
+def _move_archived_attempts_to_history(ti):
+    history = sa.table(
+        "task_instance_history",
+        sa.column("task_instance_id"),
+        *(sa.column(c) for c in _COPY_COLUMNS),
+        sa.column("dag_version_id"),
+    )
+    archived = ti.c.working_set.is_(None)
+    archived_ids = _select_archived_ids(ti)
+    op.execute(
+        history.insert().from_select(
+            ["task_instance_id", *_COPY_COLUMNS, "dag_version_id"],
+            sa.select(ti.c.id, *(ti.c[c] for c in _COPY_COLUMNS), 
ti.c.dag_version_id).where(archived),
+        )
+    )
+    hitl = sa.table("hitl_detail", sa.column("ti_id"), *(sa.column(c) for c in 
_HITL_COLUMNS))
+    hitl_history = sa.table(
+        "hitl_detail_history", sa.column("ti_history_id"), *(sa.column(c) for 
c in _HITL_COLUMNS)
+    )
+    op.execute(
+        hitl_history.insert().from_select(
+            ["ti_history_id", *_HITL_COLUMNS],
+            sa.select(hitl.c.ti_id, *(hitl.c[c] for c in _HITL_COLUMNS))
+            .join(ti, ti.c.id == hitl.c.ti_id)
+            .where(archived),
+        )
+    )
+    # SQLite runs this with foreign keys off, so cascades cannot be relied on.
+    for child in ("hitl_detail", "task_instance_note", "task_reschedule"):
+        child_table = sa.table(child, sa.column("ti_id"))
+        
op.execute(child_table.delete().where(child_table.c.ti_id.in_(archived_ids)))
+    op.execute(ti.delete().where(archived))
+
+
 def downgrade():
-    """Refuse downgrade when the predecessor cannot represent retained 
ownership."""
+    """Fold the latest attempt's data back into the legacy tables and restore 
archived attempts to history."""
     if op.get_context().as_sql:
         raise RuntimeError("Offline downgrade cannot verify retained attempt 
ownership")
     bind = op.get_bind()
-    if bind.scalar(sa.text("SELECT 1 FROM task_instance WHERE working_set IS 
NULL LIMIT 1")):
-        raise RuntimeError("Cannot downgrade attempt ownership with historical 
attempts")
-    for name in ("xcom_v2", "rtif_v2"):
-        if bind.scalar(sa.text(f"SELECT 1 FROM {name} LIMIT 1")):
-            raise RuntimeError(f"Cannot downgrade attempt ownership with data 
in {name}")
-    if bind.scalar(
-        sa.text(
-            "SELECT 1 FROM legacy_task_data_owner o JOIN task_instance t ON 
t.id=o.task_instance_id "
-            "WHERE o.dag_id<>t.dag_id OR o.task_id<>t.task_id OR 
o.run_id<>t.run_id OR o.map_index<>t.map_index LIMIT 1"
+    ti = _build_task_instance_table()
+    if _has_duplicate_attempts(bind, ti):
+        raise RuntimeError("Cannot downgrade attempt ownership with multiple 
attempts sharing coordinates")
+    if orphans := _count_orphaned_archived_attempts(bind, ti):

Review Comment:
   This refusal is reachable from a normal workflow after the upgrade, not only 
from hand-deleted rows. If a mapped task's `map_index=-1` placeholder is 
`upstream_failed` and someone clears the failed upstream plus downstream, 
`clear_task_instances` calls `prepare_db_for_next_try` on it, which archives 
the -1 row and inserts a new current -1 row. Once the upstream succeeds, 
`expand_mapped_task` rewrites that current row to `map_index=0` (or deletes it 
when index 0 already exists), and nothing moves the archived -1 row along with 
it. The downgrade then counts it as an orphan and tells the user to delete it.
   
   On 3.3 the `ON UPDATE CASCADE` / `ON DELETE CASCADE` on 
`task_instance_history_ti_fkey` carried that history row to index 0 or dropped 
it. Since 3.4 hasn't shipped, would it be simpler to have `expand_mapped_task` 
move (or delete) the archived -1 attempts together with the placeholder, so 
this guard stays a true "can't happen"? Otherwise the downgrade needs to fold 
these itself.



##########
airflow-core/tests/unit/migrations/test_0142_unify_task_attempt_ownership.py:
##########
@@ -630,43 +628,442 @@ def 
test_upgrade_preserves_uuid_children_and_legacy_cascades(populated_predecess
         assert connection.scalars(sa.select(data.c.task_instance_id)).all() == 
[]
 
 
[email protected]("unsafe_data", ["xcom_v2", "rtif_v2", "moved_owner"])
-def test_downgrade_rejects_unrepresentable_ownership(populated_predecessor, 
unsafe_data):
+COORDINATE_KEY = ("dag_id", "task_id", "run_id", "map_index")
+NEIGHBOURS = {
+    "task_id": {"task_id": "neighbour"},
+    "run_id": {"run_id": "second"},
+    "map_index": {"map_index": 0},
+    "dag_id": {"dag_id": "other"},
+}
+
+
+def insert_attempt(connection, **values) -> UUID:
+    attempt_id = uuid4()
+    row = (
+        COORDINATES | {"try_number": 1, "state": "success", "pool": 
"default_pool", "pool_slots": 1} | values
+    )
+    connection.execute(table(connection, "task_instance", 
"id").insert().values(id=attempt_id, **row))
+    return attempt_id
+
+
+def insert_dag_run(connection, dag_run_id, dag_id, run_id):
+    connection.execute(
+        table(connection, "dag_run")
+        .insert()
+        .values(
+            id=dag_run_id,
+            dag_id=dag_id,
+            run_id=run_id,
+            run_type="manual",
+            run_after=NOW,
+            state="running",
+            start_date=NOW,
+        )
+    )
+
+
+def insert_neighbour_dag_runs(connection) -> dict[tuple[str, str], int]:
+    insert_dag_run(connection, SECOND_DAG_RUN_ID, "ownership", "second")
+    insert_dag_run(connection, THIRD_DAG_RUN_ID, "other", "manual")
+    return {
+        ("ownership", "manual"): DAG_RUN_ID,
+        ("ownership", "second"): SECOND_DAG_RUN_ID,
+        ("other", "manual"): THIRD_DAG_RUN_ID,
+    }
+
+
+def insert_xcom_v2(connection, attempt_id, key, value, **values):
+    connection.execute(
+        table(connection, "xcom_v2", "id", "task_instance_id")
+        .insert()
+        .values(id=uuid4(), task_instance_id=attempt_id, key=key, value=value, 
timestamp=NOW, **values)
+    )
+
+
+def insert_rtif_v2(connection, attempt_id, rendered_fields, **values):
+    connection.execute(
+        table(connection, "rtif_v2", "id", "task_instance_id")
+        .insert()
+        .values(id=uuid4(), task_instance_id=attempt_id, 
rendered_fields=rendered_fields, **values)
+    )
+
+
+def retry_current_attempt(connection, **values) -> UUID:
+    ti = table(connection, "task_instance", "id")
+    connection.execute(
+        ti.update().where(ti.c.id == CURRENT_ID).values(working_set=None, 
archived_reason="retry")
+    )
+    return insert_attempt(connection, try_number=3, state="running", **values)
+
+
+def read_legacy_xcom(connection):
+    rows = connection.execute(table(connection, "xcom").select())
+    return {(row.dag_id, row.run_id, row.task_id, row.map_index, row.key): row 
for row in rows}
+
+
+def read_legacy_rendered_fields(connection):
+    rows = connection.execute(table(connection, 
"rendered_task_instance_fields").select())
+    return {(row.dag_id, row.run_id, row.task_id, row.map_index): row for row 
in rows}
+
+
+def drop_unique_constraint(connection, name):
+    sqlite = connection.dialect.name == "sqlite"
+    connection.commit()
+    if sqlite:
+        connection.exec_driver_sql("PRAGMA foreign_keys=OFF")
+    with (
+        Operations.context(MigrationContext.configure(connection)) as ops,
+        ops.batch_alter_table("task_instance") as batch,
+    ):
+        batch.drop_constraint(name, type_="unique")
+    connection.commit()
+    if sqlite:
+        connection.exec_driver_sql("PRAGMA foreign_keys=ON")
+
+
[email protected]("legacy_store", ["xcom", 
"rendered_task_instance_fields"])
+def 
test_downgrade_rejects_owner_that_changed_coordinates(populated_predecessor, 
legacy_store):
     connection, config = populated_predecessor
-    connection.execute(table(connection, "hitl_detail_history").delete())
-    connection.execute(table(connection, "task_instance_history").delete())
+    other_store = {"xcom", "rendered_task_instance_fields"} - {legacy_store}
+    for name in ("hitl_detail_history", "task_instance_history", *other_store):
+        connection.execute(table(connection, name).delete())
     connection.commit()
     command.upgrade(config, REVISION)
-    if unsafe_data == "moved_owner":
-        ti = table(connection, "task_instance", "id")
-        connection.execute(ti.update().where(ti.c.id == 
CURRENT_ID).values(map_index=0))
-    elif unsafe_data == "xcom_v2":
+    ti = table(connection, "task_instance", "id")
+    connection.execute(ti.update().where(ti.c.id == 
CURRENT_ID).values(map_index=0))
+    connection.commit()
+    with pytest.raises(RuntimeError, match="legacy owners changed 
coordinates"):

Review Comment:
   The removed refusal tests also asserted the schema was untouched afterwards 
(`xcom_v1` still present). Without that, moving a guard below 
`_create_history_tables()` would still pass here, and on MySQL that leaves 
`task_instance_history` behind so the next attempt fails with "table already 
exists". Could this one and the refusals at L789 and L803 keep the 
`get_table_names()` check?



##########
airflow-core/src/airflow/migrations/versions/0142_3_4_0_unify_task_attempt_ownership.py:
##########
@@ -370,24 +370,244 @@ def upgrade():
         )
 
 
+def _build_task_instance_table():
+    return sa.table(
+        "task_instance",
+        sa.column("id"),
+        *(sa.column(c) for c in _COPY_COLUMNS),
+        sa.column("working_set", sa.Boolean()),
+        sa.column("dag_version_id"),
+    )
+
+
+def _build_legacy_owner_table():
+    return sa.table(
+        "legacy_task_data_owner", *(sa.column(c) for c in _COORDINATES), 
sa.column("task_instance_id")
+    )
+
+
+def _select_latest_attempts(ti):
+    newer = ti.alias("newer")
+    # An anti-join, because grouping by max(try_number) aggregates the whole 
table however few rows need copying.
+    return (
+        sa.select(ti.c.id, *(ti.c[c] for c in _COORDINATES))
+        .where(
+            ~sa.exists().where(
+                *(newer.c[c] == ti.c[c] for c in _COORDINATES), 
newer.c.try_number > ti.c.try_number
+            )
+        )
+        .subquery()
+    )
+
+
+def _has_duplicate_attempts(bind, ti) -> bool:
+    coordinates = [ti.c[c] for c in _COORDINATES]
+    for where, group_by in (
+        (ti.c.working_set.is_not(None), coordinates),
+        (sa.true(), [*coordinates, ti.c.try_number]),
+    ):
+        duplicated = 
sa.select(1).select_from(ti).where(where).group_by(*group_by).having(sa.func.count()
 > 1)
+        if bind.scalar(duplicated.limit(1)):
+            return True
+    return False
+
+
+def _select_archived_ids(ti):
+    return sa.select(ti.c.id).where(ti.c.working_set.is_(None))
+
+
+def _count_orphaned_archived_attempts(bind, ti) -> int:
+    current = ti.alias("current")
+    orphaned = (
+        sa.select(sa.func.count())
+        .select_from(ti)
+        .where(
+            ti.c.working_set.is_(None),
+            ~sa.exists().where(
+                *(current.c[c] == ti.c[c] for c in _COORDINATES), 
current.c.working_set.is_not(None)
+            ),
+        )
+    )
+    return bind.scalar(orphaned)
+
+
+def _has_moved_owners_with_legacy_rows(bind, ti) -> bool:
+    owner = _build_legacy_owner_table()
+    moved = sa.or_(*(owner.c[c] != ti.c[c] for c in _COORDINATES))
+    for name in ("xcom_v1", "rtif_v1"):
+        legacy = sa.table(name, *(sa.column(c) for c in _COORDINATES))
+        found = (
+            sa.select(1)
+            .select_from(owner.join(ti, ti.c.id == owner.c.task_instance_id))
+            .where(moved, sa.exists().where(*(legacy.c[c] == owner.c[c] for c 
in _COORDINATES)))
+        )
+        if bind.scalar(found.limit(1)):
+            return True
+    return False
+
+
+def _delete_moved_owners(ti):
+    owner = _build_legacy_owner_table()
+    op.execute(
+        owner.delete().where(
+            sa.exists().where(
+                ti.c.id == owner.c.task_instance_id,
+                sa.or_(*(ti.c[c] != owner.c[c] for c in _COORDINATES)),
+            )
+        )
+    )
+
+
+def _delete_legacy_rows_of_archived_attempts(ti):
+    owner = _build_legacy_owner_table()
+    archived_ids = _select_archived_ids(ti)
+    # Legacy rows owned by an archived attempt are hidden from its successor, 
so keeping them would expose them again.
+    for name in ("xcom_v1", "rtif_v1"):
+        legacy = sa.table(name, *(sa.column(c) for c in _COORDINATES))
+        op.execute(
+            legacy.delete().where(
+                sa.exists().where(
+                    *(owner.c[c] == legacy.c[c] for c in _COORDINATES),
+                    owner.c.task_instance_id.in_(archived_ids),
+                )
+            )
+        )
+
+
+def _move_owners_to_current_attempts(ti):
+    owner = _build_legacy_owner_table()
+    current = ti.alias("current")
+    archived_ids = _select_archived_ids(ti)
+    xcom_v2 = sa.table("xcom_v2", sa.column("task_instance_id"))
+    rtif_v2 = sa.table("rtif_v2", sa.column("task_instance_id"))
+    # Deleting archived attempts must not cascade into legacy rows through an 
owner that points at one.
+    op.execute(
+        owner.update()
+        .where(owner.c.task_instance_id.in_(archived_ids))
+        .values(
+            task_instance_id=sa.select(current.c.id)
+            .where(*(current.c[c] == owner.c[c] for c in _COORDINATES), 
current.c.working_set.is_not(None))
+            .scalar_subquery()
+        )
+    )
+    # Legacy tables reference owners by coordinates, so attempts created after 
the upgrade need an owner row first.
+    op.execute(
+        owner.insert().from_select(
+            [*_COORDINATES, "task_instance_id"],
+            sa.select(*(ti.c[c] for c in _COORDINATES), ti.c.id).where(
+                ti.c.working_set.is_not(None),

Review Comment:
   As far as I can see no test binds this `working_set` filter. The only 
archived attempt with `xcom_v2`/`rtif_v2` rows in the tests is `HISTORY_ID`, 
which sits at a coordinate that already has an owner, so `~exists(owner)` skips 
it either way. The shape the filter protects is probably the most common one: a 
task first run after the upgrade that retried, with both tries writing XCom. 
Without the filter both attempts get an owner row for the same coordinates and 
`legacy_task_data_owner_pkey` aborts the downgrade. Could we add a test for 
that (archived try 1 and current try 2 at a post-upgrade coordinate, both 
writing `return_value` and rendered fields, only try 2 survives)?



-- 
This is an automated message from the Apache Git Service.
To respond to the message, please log on to GitHub and use the
URL above to go to the specific comment.

To unsubscribe, e-mail: [email protected]

For queries about this service, please contact Infrastructure at:
[email protected]

Reply via email to