amoghrajesh commented on code in PR #74222:
URL: https://github.com/apache/airflow/pull/74222#discussion_r4193636528


##########
airflow-core/src/airflow/models/taskinstance.py:
##########
@@ -1049,7 +1130,7 @@ def set_state(self, state: str | None, *, session: 
Session = NEW_SESSION) -> boo
         :param session: SQLAlchemy ORM Session
         :return: Was the state changed
         """
-        if self.state == state:
+        if self.state == state or (self.working_set is None and 
inspect(self).has_identity):

Review Comment:
   `set_state` on an old try now quietly returns `False`. Should it raise, so 
callers know?



##########
airflow-core/src/airflow/api_fastapi/execution_api/security.py:
##########
@@ -205,9 +222,63 @@ async def require_auth(
                 detail="Token subject does not match callback ID",
             )
 
+    if (
+        request.method not in {"GET", "HEAD", "OPTIONS"}
+        and token_scope != "callback"
+        and "connection_test_id" not in request.path_params
+        and not request.scope.get(_IN_PROCESS_NON_TI_CALLER)
+        and not request.scope.get(_REQUEST_SCOPE_LIVE_ATTEMPT_KEY)
+        and not getattr(route, _SKIP_AUTO_TI_ATTEMPT_LIVE, False)
+    ):
+        # The versions package imports routes, which depend on this module.
+        from airflow.api_fastapi.execution_api.versions.v2026_10_30 import 
IdentifyRetiredTaskStateUpdates
+
+        if IdentifyRetiredTaskStateUpdates.is_applied:

Review Comment:
   Is this intended to run for new clients only? (The if case above blocks 
writes from old tries iiuc). Task state lookups use `session.get`, which skips 
the filter btw too. So a stuck old try on a 3.3 worker can overwrite the new 
try's task state. Before this change, it got a 404.



##########
airflow-core/src/airflow/models/taskinstance.py:
##########
@@ -2814,6 +2977,40 @@ def __repr__(self):
         return prefix + f" TI ID: {self.ti_id}>"
 
 
+_CURRENT_ATTEMPTS = with_loader_criteria(
+    TaskInstance, TaskInstance.working_set.is_(True), include_aliases=True, 
propagate_to_loaders=False
+)
+
+
+def _is_primary_key_lookup(statement) -> bool:

Review Comment:
   nit: we are depending on `pk_` and `_where_criteria` (private scope). Is 
there an alternative?



##########
airflow-core/src/airflow/models/xcom.py:
##########
@@ -399,6 +424,201 @@ def deserialize_value(result: Any) -> Any:
             return result.value
 
 
+def _rows():
+    from airflow.models.dagrun import DagRun
+    from airflow.models.taskinstance import LegacyTaskDataOwner, TaskInstance
+
+    coordinates = ("dag_id", "task_id", "run_id", "map_index")
+    data = ("key", "value", "timestamp", "dag_result", "mapped_length")
+    run_join = and_(DagRun.dag_id == TaskInstance.dag_id, DagRun.run_id == 
TaskInstance.run_id)
+    context = [
+        *(getattr(TaskInstance, name) for name in coordinates),
+        DagRun.id.label("dag_run_id"),
+        DagRun.logical_date,
+        DagRun.run_after,
+    ]
+    new_rows = (
+        select(XComModelV2.task_instance_id, *(getattr(XComModelV2, name) for 
name in data), *context)
+        .select_from(XComModelV2)
+        .join(TaskInstance, XComModelV2.task_instance_id == TaskInstance.id)
+        .join(DagRun, run_join)
+    )
+    old_rows = (
+        select(LegacyTaskDataOwner.task_instance_id, *(getattr(XComModelV1, 
name) for name in data), *context)
+        .select_from(LegacyTaskDataOwner)
+        .join(
+            XComModelV1,
+            and_(*(getattr(LegacyTaskDataOwner, name) == getattr(XComModelV1, 
name) for name in coordinates)),
+        )
+        .join(TaskInstance, LegacyTaskDataOwner.task_instance_id == 
TaskInstance.id)
+        .join(DagRun, run_join)
+        .where(
+            ~select(1)
+            .select_from(XComModelV2)
+            .where(
+                XComModelV2.task_instance_id == 
LegacyTaskDataOwner.task_instance_id,
+                XComModelV2.key == XComModelV1.key,
+            )
+            .exists()
+        )
+    )
+    return union_all(new_rows, old_rows)
+
+
+def _filter_rows(rows: Subquery, *, producer_ids: Select, key: str | None) -> 
Subquery:

Review Comment:
   `_filter_rows` rewrites `union.selects` inside `cloned_traverse`. That is 
clever but a short test that pins the generated SQL would help.



##########
airflow-core/src/airflow/models/taskinstance.py:
##########
@@ -2814,6 +2977,40 @@ def __repr__(self):
         return prefix + f" TI ID: {self.ti_id}>"
 
 
+_CURRENT_ATTEMPTS = with_loader_criteria(
+    TaskInstance, TaskInstance.working_set.is_(True), include_aliases=True, 
propagate_to_loaders=False
+)
+
+
+def _is_primary_key_lookup(statement) -> bool:

Review Comment:
   Here - 
   
   `session.get(TI, id)` finds a retired try, but `select(TI).where(TI.id == 
id)` doesn't. Example: `task_state_store.py:49` vs `asset_state_store.py:64`, 
same input, different result. Could primary key lookups also hide retired rows 
by default, with `include_all_attempt` as the only way in? Then we don't need 
the `pk_` check either.



##########
airflow-core/src/airflow/api_fastapi/execution_api/security.py:
##########
@@ -205,9 +222,63 @@ async def require_auth(
                 detail="Token subject does not match callback ID",
             )
 
+    if (
+        request.method not in {"GET", "HEAD", "OPTIONS"}
+        and token_scope != "callback"
+        and "connection_test_id" not in request.path_params
+        and not request.scope.get(_IN_PROCESS_NON_TI_CALLER)
+        and not request.scope.get(_REQUEST_SCOPE_LIVE_ATTEMPT_KEY)
+        and not getattr(route, _SKIP_AUTO_TI_ATTEMPT_LIVE, False)
+    ):
+        # The versions package imports routes, which depend on this module.
+        from airflow.api_fastapi.execution_api.versions.v2026_10_30 import 
IdentifyRetiredTaskStateUpdates
+
+        if IdentifyRetiredTaskStateUpdates.is_applied:

Review Comment:
   Actually read kaxils/s comment and that makes sense too. But either ways, 
task state uses `session.get`, which skips the filter. Could 
`_get_task_scope_for_ti` reject retired tries?
   
   



##########
airflow-core/src/airflow/migrations/versions/0142_3_4_0_unify_task_attempt_ownership.py:
##########
@@ -0,0 +1,462 @@
+#
+# Licensed to the Apache Software Foundation (ASF) under one
+# or more contributor license agreements.  See the NOTICE file
+# distributed with this work for additional information
+# regarding copyright ownership.  The ASF licenses this file
+# to you under the Apache License, Version 2.0 (the
+# "License"); you may not use this file except in compliance
+# with the License.  You may obtain a copy of the License at
+#
+#   http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing,
+# software distributed under the License is distributed on an
+# "AS IS" BASIS, WITHOUT WARRANTIES OR CONDITIONS OF ANY
+# KIND, either express or implied.  See the License for the
+# specific language governing permissions and limitations
+# under the License.
+
+"""
+Unify task attempt ownership without rewriting legacy XCom data.
+
+Revision ID: e7c2a91bd540
+Revises: 90e4d18ccadf
+Create Date: 2026-09-29 12:00:00.000000
+"""
+
+from __future__ import annotations
+
+from contextlib import contextmanager
+
+import sqlalchemy as sa
+from alembic import op
+from sqlalchemy.dialects import postgresql, sqlite
+
+from airflow.migrations.db_types import TIMESTAMP, StringID
+from airflow.utils.sqlalchemy import ExecutorConfigType, ExtendedJSON, 
UtcDateTime
+
+revision = "e7c2a91bd540"
+down_revision = "90e4d18ccadf"
+branch_labels = None
+depends_on = None
+airflow_version = "3.4.0"
+
+_COORDINATES = ("dag_id", "task_id", "run_id", "map_index")
+_LEGACY = (
+    ("xcom", "xcom_v1", "xcom_task_instance_fkey", None),
+    ("rendered_task_instance_fields", "rtif_v1", "rtif_ti_fkey", None),
+)
+_COPY_COLUMNS = (
+    "task_id",
+    "dag_id",
+    "run_id",
+    "map_index",
+    "try_number",
+    "start_date",
+    "end_date",
+    "duration",
+    "state",
+    "max_tries",
+    "hostname",
+    "unixname",
+    "pool",
+    "pool_slots",
+    "queue",
+    "priority_weight",
+    "operator",
+    "custom_operator_name",
+    "queued_dttm",
+    "scheduled_dttm",
+    "queued_by_job_id",
+    "pid",
+    "executor",
+    "executor_config",
+    "updated_at",
+    "rendered_map_index",
+    "context_carrier",
+    "external_executor_id",
+    "trigger_timeout",
+    "next_method",
+    "next_kwargs",
+    "task_display_name",
+    "retry_delay_override",
+    "retry_reason",
+)
+_HITL_COLUMNS = (
+    "options",
+    "subject",
+    "body",
+    "defaults",
+    "multiple",
+    "params",
+    "assignees",
+    "created_at",
+    "responded_at",
+    "responded_by",
+    "chosen_options",
+    "params_input",
+)
+
+
+# Unlike disable_sqlite_fkeys, a failed SQLite upgrade rolls back atomically 
and foreign_keys is always restored.
+@contextmanager
+def _sqlite_rebuilds():
+    if op.get_bind().dialect.name != "sqlite":
+        yield
+        return
+    if op.get_context().as_sql:
+        raise RuntimeError("SQLite offline SQL cannot render this migration's 
table rebuilds")
+    enabled = op.get_bind().exec_driver_sql("PRAGMA foreign_keys").scalar()
+    with op.get_context().autocommit_block():
+        op.execute("PRAGMA foreign_keys=OFF")
+    try:
+        with op.get_bind().begin_nested():
+            yield
+    finally:
+        with op.get_context().autocommit_block():
+            op.execute(f"PRAGMA foreign_keys={int(enabled)}")
+
+
+def _check_source():
+    bind = op.get_bind()
+    if op.get_context().as_sql:
+        return
+    inspector = sa.inspect(bind)
+    for table, _, name, onupdate in (
+        *_LEGACY,
+        (
+            "task_instance_history",
+            None,
+            "task_instance_history_ti_fkey",
+            "CASCADE",
+        ),
+    ):
+        constraints = inspector.get_foreign_keys(table)
+        fk = next((fk for fk in constraints if fk["name"] == name), None)
+        if (
+            fk is None
+            or fk["referred_table"] != "task_instance"
+            or fk["constrained_columns"] != list(_COORDINATES)
+            or fk["referred_columns"] != list(_COORDINATES)
+            or fk["options"].get("ondelete") != "CASCADE"
+            or fk["options"].get("onupdate") != onupdate
+        ):
+            raise RuntimeError(f"Unsupported attempt ownership source schema: 
{table}.{name}")
+        if bind.dialect.name == "postgresql":
+            validated = bind.scalar(
+                sa.text(
+                    "SELECT convalidated FROM pg_constraint WHERE 
conrelid=to_regclass(:table) AND conname=:name"
+                ),
+                {"table": table, "name": name},
+            )
+            if not validated:
+                raise RuntimeError(f"Attempt ownership requires a validated 
source FK: {table}.{name}")
+
+
+def _redirect_legacy(table, constraint, target, *, onupdate=None, 
not_valid=False):
+    if op.get_bind().dialect.name == "mysql":
+        quote = op.get_bind().dialect.identifier_preparer.quote
+        columns = ", ".join(quote(column) for column in _COORDINATES)
+        onupdate_sql = f" ON UPDATE {onupdate}" if onupdate else ""
+        op.execute("SET @ti141_foreign_key_checks = 
@@SESSION.foreign_key_checks")
+        op.execute("SET SESSION foreign_key_checks = 0")
+        try:
+            op.execute(
+                f"ALTER TABLE {quote(table)} DROP FOREIGN KEY 
{quote(constraint)}, "
+                "ALGORITHM=INPLACE, LOCK=NONE"
+            )
+            op.execute(
+                f"ALTER TABLE {quote(table)} ADD CONSTRAINT 
{quote(constraint)} FOREIGN KEY ({columns}) "
+                f"REFERENCES {quote(target)} ({columns}) ON DELETE 
CASCADE{onupdate_sql}, "
+                "ALGORITHM=INPLACE, LOCK=NONE"
+            )
+        finally:
+            op.execute("SET SESSION foreign_key_checks = 
@ti141_foreign_key_checks")
+        return
+    with op.batch_alter_table(table) as batch:
+        batch.drop_constraint(constraint, type_="foreignkey")
+        batch.create_foreign_key(
+            constraint,
+            target,
+            list(_COORDINATES),
+            list(_COORDINATES),
+            ondelete="CASCADE",
+            onupdate=onupdate,
+            postgresql_not_valid=not_valid,
+        )
+
+
+def upgrade():
+    """Retain attempts and give legacy and new task data immutable UUID 
owners."""
+    _check_source()
+    with _sqlite_rebuilds():
+        owner = op.create_table(
+            "legacy_task_data_owner",
+            sa.Column("dag_id", StringID(), nullable=False),
+            sa.Column("task_id", StringID(), nullable=False),
+            sa.Column("run_id", StringID(), nullable=False),
+            sa.Column("map_index", sa.Integer(), nullable=False),
+            sa.Column("task_instance_id", sa.Uuid(), nullable=False),
+            sa.PrimaryKeyConstraint(*_COORDINATES, 
name="legacy_task_data_owner_pkey"),
+            sa.ForeignKeyConstraint(
+                ["task_instance_id"],
+                ["task_instance.id"],
+                name="legacy_task_data_owner_ti_fkey",
+                ondelete="CASCADE",
+            ),
+        )
+        op.create_index("idx_legacy_task_data_owner_ti", 
"legacy_task_data_owner", ["task_instance_id"])
+        op.add_column("log", sa.Column("task_instance_id", sa.Uuid(), 
nullable=True))
+        source = sa.table("task_instance", sa.column("id"), *(sa.column(c) for 
c in _COORDINATES))
+        op.execute(
+            owner.insert().from_select(
+                [*_COORDINATES, "task_instance_id"],
+                sa.select(*(source.c[c] for c in _COORDINATES), source.c.id),
+            )
+        )
+        for old_name, name, constraint, _ in _LEGACY:
+            op.rename_table(old_name, name)
+            # Validated source FKs plus the complete owner copy prove existing 
child ownership.
+            _redirect_legacy(
+                name,
+                constraint,
+                "legacy_task_data_owner",
+                not_valid=op.get_bind().dialect.name == "postgresql",
+            )
+        with op.batch_alter_table("task_instance_history") as batch:
+            batch.drop_constraint("task_instance_history_ti_fkey", 
type_="foreignkey")
+        with op.batch_alter_table("task_instance") as batch:
+            batch.add_column(sa.Column("working_set", sa.Boolean(), 
nullable=True, server_default=sa.true()))
+            batch.add_column(sa.Column("archived_reason", sa.String(50), 
nullable=True))
+            batch.drop_constraint("task_instance_composite_key", 
type_="unique")
+            batch.create_unique_constraint("task_instance_current_key", 
[*_COORDINATES, "working_set"])
+            batch.create_unique_constraint("task_instance_try_key", 
[*_COORDINATES, "try_number"])
+            batch.create_check_constraint(
+                "ti_working_set_true_or_null", "working_set IS NULL OR 
working_set = TRUE"
+            )
+        ti = sa.table(
+            "task_instance",
+            sa.column("id"),
+            *(sa.column(c) for c in _COPY_COLUMNS),
+            sa.column("working_set"),
+            sa.column("archived_reason"),
+            sa.column("dag_version_id"),
+        )
+        history = sa.table(
+            "task_instance_history",
+            sa.column("task_instance_id"),
+            *(sa.column(c) for c in _COPY_COLUMNS),
+            sa.column("dag_version_id"),
+        )
+        historical_max_tries = sa.func.coalesce(
+            history.c.max_tries,
+            sa.case((history.c.try_number > 0, history.c.try_number - 1), 
else_=0),
+        )
+        version = sa.table("dag_version", sa.column("id"))
+        history_rows = sa.select(
+            history.c.task_instance_id,
+            *(historical_max_tries if c == "max_tries" else history.c[c] for c 
in _COPY_COLUMNS),
+            sa.null(),
+            sa.literal("legacy"),
+            version.c.id,
+        ).select_from(history.outerjoin(version, history.c.dag_version_id == 
version.c.id))
+        dialect = op.get_bind().dialect.name
+        if dialect == "mysql":
+            existing = ti.alias("existing")
+            history_rows = history_rows.where(
+                ~sa.select(1)
+                .select_from(existing)
+                .where(*(existing.c[c] == history.c[c] for c in 
(*_COORDINATES, "try_number")))
+                .exists()
+            )
+        columns = [
+            "id",
+            *_COPY_COLUMNS,
+            "working_set",
+            "archived_reason",
+            "dag_version_id",
+        ]
+        if dialect == "postgresql":
+            insert = (
+                postgresql.insert(ti)
+                .from_select(columns, history_rows)
+                .on_conflict_do_nothing(constraint="task_instance_try_key")
+            )
+        elif dialect == "sqlite":
+            insert = (
+                sqlite.insert(ti)
+                .from_select(columns, history_rows.where(sa.true()))
+                .on_conflict_do_nothing(index_elements=[*_COORDINATES, 
"try_number"])
+            )
+        else:
+            insert = ti.insert().from_select(columns, history_rows)
+        op.execute(insert)
+        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.insert().from_select(
+                ["ti_id", *_HITL_COLUMNS],
+                sa.select(hitl_history.c.ti_history_id, *(hitl_history.c[c] 
for c in _HITL_COLUMNS)).join(
+                    ti,
+                    sa.and_(ti.c.id == hitl_history.c.ti_history_id, 
ti.c.working_set.is_(None)),
+                ),
+            )
+        )
+        op.drop_table("hitl_detail_history")
+        op.drop_table("task_instance_history")
+        for name, columns in (
+            ("ti_current_state", ["working_set", "state"]),
+            ("ti_current_dag_run", ["working_set", "dag_id", "run_id", 
"state"]),
+        ):
+            op.create_index(
+                name,
+                "task_instance",
+                columns,
+                postgresql_where=sa.text("working_set IS TRUE"),
+                sqlite_where=sa.text("working_set IS TRUE"),
+            )
+        op.create_table(
+            "xcom_v2",
+            sa.Column("id", sa.Uuid(), nullable=False),
+            sa.Column("task_instance_id", sa.Uuid(), nullable=False),
+            sa.Column("key", StringID(length=512), nullable=False),
+            sa.Column("value", sa.JSON().with_variant(postgresql.JSONB(), 
"postgresql")),
+            sa.Column("timestamp", TIMESTAMP(), nullable=False),
+            sa.Column("dag_result", sa.Boolean(), nullable=True),
+            sa.Column("mapped_length", sa.Integer(), nullable=True),
+            sa.PrimaryKeyConstraint("id", name="xcom_v2_pkey"),
+            sa.UniqueConstraint("task_instance_id", "key", 
name="xcom_v2_ti_key_uq"),
+            sa.ForeignKeyConstraint(
+                ["task_instance_id"], ["task_instance.id"], 
name="xcom_v2_ti_fkey", ondelete="CASCADE"
+            ),
+            sa.CheckConstraint("mapped_length >= 0", 
name="xcom_v2_mapped_length_not_negative"),
+        )
+        op.create_table(
+            "rtif_v2",
+            sa.Column("id", sa.Uuid(), nullable=False),
+            sa.Column("task_instance_id", sa.Uuid(), nullable=False),
+            sa.Column("rendered_fields", sa.JSON(), nullable=False),
+            sa.Column("k8s_pod_yaml", sa.JSON(), nullable=True),
+            sa.PrimaryKeyConstraint("id", name="rtif_v2_pkey"),
+            sa.UniqueConstraint("task_instance_id", name="rtif_v2_ti_uq"),
+            sa.ForeignKeyConstraint(
+                ["task_instance_id"], ["task_instance.id"], 
name="rtif_v2_ti_fkey", ondelete="CASCADE"
+            ),
+        )
+
+
+def downgrade():
+    """Refuse downgrade when the predecessor cannot represent retained 
ownership."""
+    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")):

Review Comment:
   Yeah, for me blocking `>>` dropping data.  But, in this case, downgrade 
could move old tries back into TIH right? The only data lost would be xcom and 
rendered fields of old tries, and 3.3 deletes those anyway I think. So a 
downgrade with no real loss looks doable?



##########
airflow-core/src/airflow/migrations/versions/0142_3_4_0_unify_task_attempt_ownership.py:
##########
@@ -0,0 +1,462 @@
+#
+# Licensed to the Apache Software Foundation (ASF) under one
+# or more contributor license agreements.  See the NOTICE file
+# distributed with this work for additional information
+# regarding copyright ownership.  The ASF licenses this file
+# to you under the Apache License, Version 2.0 (the
+# "License"); you may not use this file except in compliance
+# with the License.  You may obtain a copy of the License at
+#
+#   http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing,
+# software distributed under the License is distributed on an
+# "AS IS" BASIS, WITHOUT WARRANTIES OR CONDITIONS OF ANY
+# KIND, either express or implied.  See the License for the
+# specific language governing permissions and limitations
+# under the License.
+
+"""
+Unify task attempt ownership without rewriting legacy XCom data.
+
+Revision ID: e7c2a91bd540
+Revises: 90e4d18ccadf
+Create Date: 2026-09-29 12:00:00.000000
+"""
+
+from __future__ import annotations
+
+from contextlib import contextmanager
+
+import sqlalchemy as sa
+from alembic import op
+from sqlalchemy.dialects import postgresql, sqlite
+
+from airflow.migrations.db_types import TIMESTAMP, StringID
+from airflow.utils.sqlalchemy import ExecutorConfigType, ExtendedJSON, 
UtcDateTime
+
+revision = "e7c2a91bd540"
+down_revision = "90e4d18ccadf"
+branch_labels = None
+depends_on = None
+airflow_version = "3.4.0"
+
+_COORDINATES = ("dag_id", "task_id", "run_id", "map_index")
+_LEGACY = (
+    ("xcom", "xcom_v1", "xcom_task_instance_fkey", None),
+    ("rendered_task_instance_fields", "rtif_v1", "rtif_ti_fkey", None),
+)
+_COPY_COLUMNS = (
+    "task_id",
+    "dag_id",
+    "run_id",
+    "map_index",
+    "try_number",
+    "start_date",
+    "end_date",
+    "duration",
+    "state",
+    "max_tries",
+    "hostname",
+    "unixname",
+    "pool",
+    "pool_slots",
+    "queue",
+    "priority_weight",
+    "operator",
+    "custom_operator_name",
+    "queued_dttm",
+    "scheduled_dttm",
+    "queued_by_job_id",
+    "pid",
+    "executor",
+    "executor_config",
+    "updated_at",
+    "rendered_map_index",
+    "context_carrier",
+    "external_executor_id",
+    "trigger_timeout",
+    "next_method",
+    "next_kwargs",
+    "task_display_name",
+    "retry_delay_override",
+    "retry_reason",
+)
+_HITL_COLUMNS = (
+    "options",
+    "subject",
+    "body",
+    "defaults",
+    "multiple",
+    "params",
+    "assignees",
+    "created_at",
+    "responded_at",
+    "responded_by",
+    "chosen_options",
+    "params_input",
+)
+
+
+# Unlike disable_sqlite_fkeys, a failed SQLite upgrade rolls back atomically 
and foreign_keys is always restored.
+@contextmanager
+def _sqlite_rebuilds():
+    if op.get_bind().dialect.name != "sqlite":
+        yield
+        return
+    if op.get_context().as_sql:
+        raise RuntimeError("SQLite offline SQL cannot render this migration's 
table rebuilds")
+    enabled = op.get_bind().exec_driver_sql("PRAGMA foreign_keys").scalar()
+    with op.get_context().autocommit_block():
+        op.execute("PRAGMA foreign_keys=OFF")
+    try:
+        with op.get_bind().begin_nested():
+            yield
+    finally:
+        with op.get_context().autocommit_block():
+            op.execute(f"PRAGMA foreign_keys={int(enabled)}")
+
+
+def _check_source():
+    bind = op.get_bind()
+    if op.get_context().as_sql:
+        return
+    inspector = sa.inspect(bind)
+    for table, _, name, onupdate in (
+        *_LEGACY,
+        (
+            "task_instance_history",
+            None,
+            "task_instance_history_ti_fkey",
+            "CASCADE",
+        ),
+    ):
+        constraints = inspector.get_foreign_keys(table)
+        fk = next((fk for fk in constraints if fk["name"] == name), None)
+        if (
+            fk is None
+            or fk["referred_table"] != "task_instance"
+            or fk["constrained_columns"] != list(_COORDINATES)
+            or fk["referred_columns"] != list(_COORDINATES)
+            or fk["options"].get("ondelete") != "CASCADE"
+            or fk["options"].get("onupdate") != onupdate
+        ):
+            raise RuntimeError(f"Unsupported attempt ownership source schema: 
{table}.{name}")
+        if bind.dialect.name == "postgresql":
+            validated = bind.scalar(
+                sa.text(
+                    "SELECT convalidated FROM pg_constraint WHERE 
conrelid=to_regclass(:table) AND conname=:name"
+                ),
+                {"table": table, "name": name},
+            )
+            if not validated:
+                raise RuntimeError(f"Attempt ownership requires a validated 
source FK: {table}.{name}")
+
+
+def _redirect_legacy(table, constraint, target, *, onupdate=None, 
not_valid=False):
+    if op.get_bind().dialect.name == "mysql":
+        quote = op.get_bind().dialect.identifier_preparer.quote
+        columns = ", ".join(quote(column) for column in _COORDINATES)
+        onupdate_sql = f" ON UPDATE {onupdate}" if onupdate else ""
+        op.execute("SET @ti141_foreign_key_checks = 
@@SESSION.foreign_key_checks")
+        op.execute("SET SESSION foreign_key_checks = 0")
+        try:
+            op.execute(
+                f"ALTER TABLE {quote(table)} DROP FOREIGN KEY 
{quote(constraint)}, "
+                "ALGORITHM=INPLACE, LOCK=NONE"
+            )
+            op.execute(
+                f"ALTER TABLE {quote(table)} ADD CONSTRAINT 
{quote(constraint)} FOREIGN KEY ({columns}) "
+                f"REFERENCES {quote(target)} ({columns}) ON DELETE 
CASCADE{onupdate_sql}, "
+                "ALGORITHM=INPLACE, LOCK=NONE"
+            )
+        finally:
+            op.execute("SET SESSION foreign_key_checks = 
@ti141_foreign_key_checks")
+        return
+    with op.batch_alter_table(table) as batch:
+        batch.drop_constraint(constraint, type_="foreignkey")
+        batch.create_foreign_key(
+            constraint,
+            target,
+            list(_COORDINATES),
+            list(_COORDINATES),
+            ondelete="CASCADE",
+            onupdate=onupdate,
+            postgresql_not_valid=not_valid,
+        )
+
+
+def upgrade():
+    """Retain attempts and give legacy and new task data immutable UUID 
owners."""
+    _check_source()
+    with _sqlite_rebuilds():
+        owner = op.create_table(
+            "legacy_task_data_owner",
+            sa.Column("dag_id", StringID(), nullable=False),
+            sa.Column("task_id", StringID(), nullable=False),
+            sa.Column("run_id", StringID(), nullable=False),
+            sa.Column("map_index", sa.Integer(), nullable=False),
+            sa.Column("task_instance_id", sa.Uuid(), nullable=False),
+            sa.PrimaryKeyConstraint(*_COORDINATES, 
name="legacy_task_data_owner_pkey"),
+            sa.ForeignKeyConstraint(
+                ["task_instance_id"],
+                ["task_instance.id"],
+                name="legacy_task_data_owner_ti_fkey",
+                ondelete="CASCADE",
+            ),
+        )
+        op.create_index("idx_legacy_task_data_owner_ti", 
"legacy_task_data_owner", ["task_instance_id"])
+        op.add_column("log", sa.Column("task_instance_id", sa.Uuid(), 
nullable=True))
+        source = sa.table("task_instance", sa.column("id"), *(sa.column(c) for 
c in _COORDINATES))
+        op.execute(
+            owner.insert().from_select(
+                [*_COORDINATES, "task_instance_id"],
+                sa.select(*(source.c[c] for c in _COORDINATES), source.c.id),
+            )
+        )
+        for old_name, name, constraint, _ in _LEGACY:
+            op.rename_table(old_name, name)
+            # Validated source FKs plus the complete owner copy prove existing 
child ownership.
+            _redirect_legacy(
+                name,
+                constraint,
+                "legacy_task_data_owner",
+                not_valid=op.get_bind().dialect.name == "postgresql",
+            )
+        with op.batch_alter_table("task_instance_history") as batch:

Review Comment:
   PG looks fine, mysql is slow and `CHECK` alone is 211 ms..should we just 
skip the `CHECK` on mysql? We never write `FALSE` so check is guarding nothing 
actually.



##########
airflow-core/src/airflow/models/taskinstance.py:
##########
@@ -1071,28 +1122,98 @@ def is_premature(self) -> bool:
         # is the task still in the retry waiting period?
         return self.state == TaskInstanceState.UP_FOR_RETRY and not 
self.ready_for_retry()
 
-    def prepare_db_for_next_try(self, session: Session):
-        """Archive this attempt and allocate the next attempt's UUID and try 
number."""
-        from airflow.models.taskinstancehistory import TaskInstanceHistory
+    def retire(self, *, reason: str, session: Session) -> None:
+        """Remove this attempt from the working set while retaining its UUID 
and children."""
+        current = session.scalar(
+            select(TaskInstance.working_set).where(TaskInstance.id == 
self.id).with_for_update()
+        )
+        if current is not True:
+            raise ValueError("A retired task instance cannot be retired again")
+        if self.state not in State.finished:
+            self.state = TaskInstanceState.FAILED
+            if self.end_date is None:
+                self.end_date = timezone.utcnow()
+                self.set_duration()
+        self.working_set = None
+        self.archived_reason = reason
+        self.trigger_id = None
+        session.flush()
 
-        TaskInstanceHistory.record_ti(self, session=session)
-        session.execute(delete(TaskReschedule).filter_by(ti_id=self.id))
-        self.external_executor_id = None
-        self.id = uuid7()
-        self.try_number += 1
+    @classmethod
+    def delete_attempts(
+        cls,
+        *,
+        dag_id: str,
+        run_id: str,
+        task_id: str,
+        map_index: int | None = None,
+        session: Session,
+    ) -> None:
+        """Delete every attempt, current and retired, of a task; all map 
indexes if ``map_index`` is None."""
+        statement = delete(cls).where(cls.dag_id == dag_id, cls.run_id == 
run_id, cls.task_id == task_id)
+        if map_index is not None:
+            statement = statement.where(cls.map_index == map_index)
+        session.execute(statement)
+
+    @classmethod
+    def get_last_try_numbers(
+        cls,
+        *,
+        dag_id: str,
+        task_id: str,
+        run_id: str,
+        map_indexes: Collection[int] | None = None,
+        session: Session,
+    ) -> dict[int, int]:
+        """Return the highest try number per map index across current and 
historical task instances."""
+        if map_indexes is not None and not map_indexes:
+            return {}
+        statement = (
+            select(cls.map_index, func.max(cls.try_number))
+            .where(cls.dag_id == dag_id, cls.task_id == task_id, cls.run_id == 
run_id)
+            .group_by(cls.map_index)
+        )
+        if map_indexes is not None:
+            statement = statement.where(cls.map_index.in_(map_indexes))
+        return {map_index: last_try for map_index, last_try in 
session.execute(statement)}
+
+    def prepare_db_for_next_try(self, session: Session) -> TaskInstance:

Review Comment:
   Follow up is fine



-- 
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