This is an automated email from the ASF dual-hosted git repository.

amoghdesai pushed a commit to branch main
in repository https://gitbox.apache.org/repos/asf/airflow.git


The following commit(s) were added to refs/heads/main by this push:
     new bb77ebf1a23 AIP-72: Port Registering of Asset Changes to Task SDK on 
task completion (#45924)
bb77ebf1a23 is described below

commit bb77ebf1a2384169c4cd3132b36c54d779473886
Author: Amogh Desai <[email protected]>
AuthorDate: Fri Jan 24 12:11:38 2025 +0530

    AIP-72: Port Registering of Asset Changes to Task SDK on task completion 
(#45924)
---
 .../api_fastapi/execution_api/datamodels/asset.py  |  15 +++
 .../execution_api/datamodels/taskinstance.py       |  48 ++++++++-
 .../execution_api/routes/task_instances.py         |  14 ++-
 airflow/models/taskinstance.py                     |  82 +++++++++------
 task_sdk/src/airflow/sdk/api/client.py             |   6 ++
 .../src/airflow/sdk/api/datamodels/_generated.py   |  27 ++++-
 task_sdk/src/airflow/sdk/execution_time/comms.py   |  15 ++-
 .../src/airflow/sdk/execution_time/supervisor.py   |  15 ++-
 .../src/airflow/sdk/execution_time/task_runner.py  |  42 +++++++-
 task_sdk/tests/execution_time/test_supervisor.py   |  17 +++
 task_sdk/tests/execution_time/test_task_runner.py  |  71 ++++++++++++-
 .../execution_api/routes/test_task_instances.py    | 116 +++++++++++++++++++--
 12 files changed, 416 insertions(+), 52 deletions(-)

diff --git a/airflow/api_fastapi/execution_api/datamodels/asset.py 
b/airflow/api_fastapi/execution_api/datamodels/asset.py
index 6d3a53c3e4c..29b260c291c 100644
--- a/airflow/api_fastapi/execution_api/datamodels/asset.py
+++ b/airflow/api_fastapi/execution_api/datamodels/asset.py
@@ -34,3 +34,18 @@ class AssetAliasResponse(BaseModel):
 
     name: str
     group: str
+
+
+class AssetProfile(BaseModel):
+    """
+    Profile of an Asset.
+
+    Asset will have name, uri and asset_type defined.
+    AssetNameRef will have name and asset_type defined.
+    AssetUriRef will have uri and asset_type defined.
+
+    """
+
+    name: str | None = None
+    uri: str | None = None
+    asset_type: str
diff --git a/airflow/api_fastapi/execution_api/datamodels/taskinstance.py 
b/airflow/api_fastapi/execution_api/datamodels/taskinstance.py
index 5e8c267b82f..6cc82259cf7 100644
--- a/airflow/api_fastapi/execution_api/datamodels/taskinstance.py
+++ b/airflow/api_fastapi/execution_api/datamodels/taskinstance.py
@@ -21,10 +21,19 @@ import uuid
 from datetime import timedelta
 from typing import Annotated, Any, Literal, Union
 
-from pydantic import AwareDatetime, Discriminator, Field, Tag, TypeAdapter, 
WithJsonSchema, field_validator
+from pydantic import (
+    AwareDatetime,
+    Discriminator,
+    Field,
+    Tag,
+    TypeAdapter,
+    WithJsonSchema,
+    field_validator,
+)
 
 from airflow.api_fastapi.common.types import UtcDateTime
 from airflow.api_fastapi.core_api.base import BaseModel
+from airflow.api_fastapi.execution_api.datamodels.asset import AssetProfile
 from airflow.api_fastapi.execution_api.datamodels.connection import 
ConnectionResponse
 from airflow.api_fastapi.execution_api.datamodels.variable import 
VariableResponse
 from airflow.utils.state import IntermediateTIState, TaskInstanceState as 
TIState, TerminalTIState
@@ -52,14 +61,41 @@ class TIEnterRunningPayload(BaseModel):
 
 
 class TITerminalStatePayload(BaseModel):
-    """Schema for updating TaskInstance to a terminal state (e.g., SUCCESS or 
FAILED)."""
+    """Schema for updating TaskInstance to a terminal state except SUCCESS 
state."""
 
-    state: TerminalTIState
+    state: Literal[
+        TerminalTIState.FAILED,
+        TerminalTIState.SKIPPED,
+        TerminalTIState.REMOVED,
+        TerminalTIState.FAIL_WITHOUT_RETRY,
+    ]
 
     end_date: UtcDateTime
     """When the task completed executing"""
 
 
+class TISuccessStatePayload(BaseModel):
+    """Schema for updating TaskInstance to success state."""
+
+    state: Annotated[
+        Literal[TerminalTIState.SUCCESS],
+        # Specify a default in the schema, but not in code, so Pydantic marks 
it as required.
+        WithJsonSchema(
+            {
+                "type": "string",
+                "enum": [TerminalTIState.SUCCESS],
+                "default": TerminalTIState.SUCCESS,
+            }
+        ),
+    ]
+
+    end_date: UtcDateTime
+    """When the task completed executing"""
+
+    task_outlets: Annotated[list[AssetProfile], Field(default_factory=list)]
+    outlet_events: Annotated[list[Any], Field(default_factory=list)]
+
+
 class TITargetStatePayload(BaseModel):
     """Schema for updating TaskInstance to a target state, excluding terminal 
and running states."""
 
@@ -123,7 +159,10 @@ def ti_state_discriminator(v: dict[str, str] | BaseModel) 
-> str:
         state = v.get("state")
     else:
         state = getattr(v, "state", None)
-    if state in set(TerminalTIState):
+
+    if state == TIState.SUCCESS:
+        return "success"
+    elif state in set(TerminalTIState):
         return "_terminal_"
     elif state == TIState.DEFERRED:
         return "deferred"
@@ -137,6 +176,7 @@ def ti_state_discriminator(v: dict[str, str] | BaseModel) 
-> str:
 TIStateUpdate = Annotated[
     Union[
         Annotated[TITerminalStatePayload, Tag("_terminal_")],
+        Annotated[TISuccessStatePayload, Tag("success")],
         Annotated[TITargetStatePayload, Tag("_other_")],
         Annotated[TIDeferredStatePayload, Tag("deferred")],
         Annotated[TIRescheduleStatePayload, Tag("up_for_reschedule")],
diff --git a/airflow/api_fastapi/execution_api/routes/task_instances.py 
b/airflow/api_fastapi/execution_api/routes/task_instances.py
index 4899e93c612..899017e612d 100644
--- a/airflow/api_fastapi/execution_api/routes/task_instances.py
+++ b/airflow/api_fastapi/execution_api/routes/task_instances.py
@@ -38,6 +38,7 @@ from 
airflow.api_fastapi.execution_api.datamodels.taskinstance import (
     TIRescheduleStatePayload,
     TIRunContext,
     TIStateUpdate,
+    TISuccessStatePayload,
     TITerminalStatePayload,
 )
 from airflow.models.dagrun import DagRun as DR
@@ -226,7 +227,7 @@ def ti_update_state(
         )
 
     # We exclude_unset to avoid updating fields that are not set in the payload
-    data = ti_patch_payload.model_dump(exclude_unset=True)
+    data = ti_patch_payload.model_dump(exclude={"task_outlets", 
"outlet_events"}, exclude_unset=True)
 
     query = update(TI).where(TI.id == ti_id_str).values(data)
 
@@ -243,6 +244,17 @@ def ti_update_state(
             else:
                 updated_state = State.FAILED
         query = query.values(state=updated_state)
+    elif isinstance(ti_patch_payload, TISuccessStatePayload):
+        query = TI.duration_expression_update(ti_patch_payload.end_date, 
query, session.bind)
+        updated_state = ti_patch_payload.state
+        task_instance = session.get(TI, ti_id_str)
+        TI.register_asset_changes_in_db(
+            task_instance,
+            ti_patch_payload.task_outlets,  # type: ignore
+            ti_patch_payload.outlet_events,
+            session,
+        )
+        query = query.values(state=updated_state)
     elif isinstance(ti_patch_payload, TIDeferredStatePayload):
         # Calculate timeout if it was passed
         timeout = None
diff --git a/airflow/models/taskinstance.py b/airflow/models/taskinstance.py
index c0f6d20703b..f1aa3a8236e 100644
--- a/airflow/models/taskinstance.py
+++ b/airflow/models/taskinstance.py
@@ -105,6 +105,7 @@ from airflow.models.taskmap import TaskMap
 from airflow.models.taskreschedule import TaskReschedule
 from airflow.models.xcom import LazyXComSelectSequence, XCom
 from airflow.plugins_manager import integrate_macros_plugins
+from airflow.sdk.api.datamodels._generated import AssetProfile
 from airflow.sdk.definitions._internal.templater import SandboxedEnvironment
 from airflow.sdk.definitions.asset import Asset, AssetAlias, AssetNameRef, 
AssetUniqueKey, AssetUriRef
 from airflow.sdk.definitions.taskgroup import MappedTaskGroup
@@ -160,7 +161,7 @@ if TYPE_CHECKING:
     from airflow.models.dagrun import DagRun
     from airflow.sdk.definitions._internal.abstractoperator import Operator
     from airflow.sdk.definitions.dag import DAG
-    from airflow.sdk.types import OutletEventAccessorsProtocol, 
RuntimeTaskInstanceProtocol
+    from airflow.sdk.types import RuntimeTaskInstanceProtocol
     from airflow.typing_compat import Literal, TypeGuard
     from airflow.utils.task_group import TaskGroup
 
@@ -352,7 +353,29 @@ def _run_raw_task(
         if not test_mode:
             _add_log(event=ti.state, task_instance=ti, session=session)
             if ti.state == TaskInstanceState.SUCCESS:
-                ti._register_asset_changes(events=context["outlet_events"], 
session=session)
+                added_alias_to_task_outlet = False
+                task_outlets = []
+                outlet_events = []
+                events = context["outlet_events"]
+                for obj in ti.task.outlets or []:
+                    # Lineage can have other types of objects besides assets
+                    asset_type = type(obj).__name__
+                    if isinstance(obj, Asset):
+                        task_outlets.append(AssetProfile(name=obj.name, 
uri=obj.uri, asset_type=asset_type))
+                        outlet_events.append(attrs.asdict(events[obj]))  # 
type: ignore
+                    elif isinstance(obj, AssetNameRef):
+                        task_outlets.append(AssetProfile(name=obj.name, 
asset_type=asset_type))
+                        outlet_events.append(attrs.asdict(events))  # type: 
ignore
+                    elif isinstance(obj, AssetUriRef):
+                        task_outlets.append(AssetProfile(uri=obj.uri, 
asset_type=asset_type))
+                        outlet_events.append(attrs.asdict(events))  # type: 
ignore
+                    elif isinstance(obj, AssetAlias):
+                        if not added_alias_to_task_outlet:
+                            
task_outlets.append(AssetProfile(asset_type=asset_type))
+                            added_alias_to_task_outlet = True
+                        for asset_alias_event in 
events[obj].asset_alias_events:
+                            
outlet_events.append(attrs.asdict(asset_alias_event))
+                TaskInstance.register_asset_changes_in_db(ti, task_outlets, 
outlet_events, session=session)
 
             TaskInstance.save_to_db(ti=ti, session=session)
             if ti.state == TaskInstanceState.SUCCESS:
@@ -2733,49 +2756,46 @@ class TaskInstance(Base, LoggingMixin):
             session=session,
         )
 
-    def _register_asset_changes(
-        self, *, events: OutletEventAccessorsProtocol, session: Session | None 
= None
-    ) -> None:
-        if session:
-            TaskInstance._register_asset_changes_int(ti=self, events=events, 
session=session)
-        else:
-            TaskInstance._register_asset_changes_int(ti=self, events=events)
-
     @staticmethod
     @provide_session
-    def _register_asset_changes_int(
-        ti: TaskInstance, *, events: OutletEventAccessorsProtocol, session: 
Session = NEW_SESSION
+    def register_asset_changes_in_db(
+        ti: TaskInstance,
+        task_outlets: list[AssetProfile],
+        outlet_events: list[Any],
+        session: Session = NEW_SESSION,
     ) -> None:
-        if TYPE_CHECKING:
-            assert ti.task
-
         # One task only triggers one asset event for each asset with the same 
extra.
         # This tuple[asset uri, extra] to sets alias names mapping is used to 
find whether
         # there're assets with same uri but different extra that we need to 
emit more than one asset events.
         asset_alias_names: dict[tuple[AssetUniqueKey, frozenset], set[str]] = 
defaultdict(set)
-
         asset_name_refs: set[str] = set()
         asset_uri_refs: set[str] = set()
 
-        for obj in ti.task.outlets or []:
+        for obj in task_outlets:
             ti.log.debug("outlet obj %s", obj)
             # Lineage can have other types of objects besides assets
-            if isinstance(obj, Asset):
+            if obj.asset_type == Asset.__name__:
                 asset_manager.register_asset_change(
                     task_instance=ti,
-                    asset=obj,
-                    extra=events[obj].extra,
+                    asset=Asset(name=obj.name, uri=obj.uri),  # type: ignore
+                    extra=outlet_events[0]["extra"],
                     session=session,
                 )
-            elif isinstance(obj, AssetNameRef):
-                asset_name_refs.add(obj.name)
-            elif isinstance(obj, AssetUriRef):
-                asset_uri_refs.add(obj.uri)
-            elif isinstance(obj, AssetAlias):
-                for asset_alias_event in events[obj].asset_alias_events:
-                    asset_alias_name = asset_alias_event.source_alias_name
-                    asset_unique_key = asset_alias_event.dest_asset_key
-                    frozen_extra = frozenset(asset_alias_event.extra.items())
+            elif obj.asset_type == AssetNameRef.__name__:
+                asset_name_refs.add(obj.name)  # type: ignore
+            elif obj.asset_type == AssetUriRef.__name__:
+                asset_uri_refs.add(obj.uri)  # type: ignore
+            elif obj.asset_type == AssetAlias.__name__:
+                outlet_events = list(
+                    map(
+                        lambda event: {**event, "dest_asset_key": 
AssetUniqueKey(**event["dest_asset_key"])},
+                        outlet_events,
+                    )
+                )
+                for asset_alias_event in outlet_events:
+                    asset_alias_name = asset_alias_event["source_alias_name"]
+                    asset_unique_key = asset_alias_event["dest_asset_key"]
+                    frozen_extra = 
frozenset(asset_alias_event["extra"].items())
                     asset_alias_names[(asset_unique_key, 
frozen_extra)].add(asset_alias_name)
 
         asset_unique_keys = {key for key, _ in asset_alias_names}
@@ -2827,7 +2847,7 @@ class TaskInstance(Base, LoggingMixin):
             asset_manager.register_asset_change(
                 task_instance=ti,
                 asset=asset_model,
-                extra=events[asset_model].extra,
+                extra=outlet_events[asset_model].extra,
                 session=session,
             )
         asset_stmt = 
select(AssetModel).where(AssetModel.uri.in_(asset_uri_refs), 
AssetModel.active.has())
@@ -2836,7 +2856,7 @@ class TaskInstance(Base, LoggingMixin):
             asset_manager.register_asset_change(
                 task_instance=ti,
                 asset=asset_model,
-                extra=events[asset_model].extra,
+                extra=outlet_events[asset_model].extra,
                 session=session,
             )
 
diff --git a/task_sdk/src/airflow/sdk/api/client.py 
b/task_sdk/src/airflow/sdk/api/client.py
index b984669aa74..443256e3a67 100644
--- a/task_sdk/src/airflow/sdk/api/client.py
+++ b/task_sdk/src/airflow/sdk/api/client.py
@@ -44,6 +44,7 @@ from airflow.sdk.api.datamodels._generated import (
     TIHeartbeatInfo,
     TIRescheduleStatePayload,
     TIRunContext,
+    TISuccessStatePayload,
     TITerminalStatePayload,
     ValidationError as RemoteValidationError,
     VariablePostBody,
@@ -136,6 +137,11 @@ class TaskInstanceOperations:
         body = TITerminalStatePayload(end_date=when, 
state=TerminalTIState(state))
         self.client.patch(f"task-instances/{id}/state", 
content=body.model_dump_json())
 
+    def succeed(self, id: uuid.UUID, when: datetime, task_outlets, 
outlet_events):
+        """Tell the API server that this TI has succeeded."""
+        body = TISuccessStatePayload(end_date=when, task_outlets=task_outlets, 
outlet_events=outlet_events)
+        self.client.patch(f"task-instances/{id}/state", 
content=body.model_dump_json())
+
     def heartbeat(self, id: uuid.UUID, pid: int):
         body = TIHeartbeatInfo(pid=pid, hostname=get_hostname())
         self.client.put(f"task-instances/{id}/heartbeat", 
content=body.model_dump_json())
diff --git a/task_sdk/src/airflow/sdk/api/datamodels/_generated.py 
b/task_sdk/src/airflow/sdk/api/datamodels/_generated.py
index d91ecf841d9..3383e61d1c3 100644
--- a/task_sdk/src/airflow/sdk/api/datamodels/_generated.py
+++ b/task_sdk/src/airflow/sdk/api/datamodels/_generated.py
@@ -29,6 +29,20 @@ from uuid import UUID
 from pydantic import BaseModel, ConfigDict, Field
 
 
+class AssetProfile(BaseModel):
+    """
+    Profile of an Asset.
+
+    Asset will have name, uri and asset_type defined.
+    AssetNameRef will have name and asset_type defined.
+    AssetUriRef will have uri and asset_type defined.
+    """
+
+    name: Annotated[str | None, Field(title="Name")] = None
+    uri: Annotated[str | None, Field(title="Uri")] = None
+    asset_type: Annotated[str, Field(title="Asset Type")]
+
+
 class AssetResponse(BaseModel):
     """
     Asset schema for responses with fields that are needed for Runtime.
@@ -134,6 +148,17 @@ class TIRescheduleStatePayload(BaseModel):
     end_date: Annotated[datetime, Field(title="End Date")]
 
 
+class TISuccessStatePayload(BaseModel):
+    """
+    Schema for updating TaskInstance to success state.
+    """
+
+    state: Annotated[Literal["success"] | None, Field(title="State")] = 
"success"
+    end_date: Annotated[datetime, Field(title="End Date")]
+    task_outlets: Annotated[list[AssetProfile] | None, Field(title="Task 
Outlets")] = None
+    outlet_events: Annotated[list | None, Field(title="Outlet Events")] = None
+
+
 class TITargetStatePayload(BaseModel):
     """
     Schema for updating TaskInstance to a target state, excluding terminal and 
running states.
@@ -243,7 +268,7 @@ class TIRunContext(BaseModel):
 
 class TITerminalStatePayload(BaseModel):
     """
-    Schema for updating TaskInstance to a terminal state (e.g., SUCCESS or 
FAILED).
+    Schema for updating TaskInstance to a terminal state except SUCCESS state.
     """
 
     state: TerminalTIState
diff --git a/task_sdk/src/airflow/sdk/execution_time/comms.py 
b/task_sdk/src/airflow/sdk/execution_time/comms.py
index 007e3fe10fe..3ab8addc8bb 100644
--- a/task_sdk/src/airflow/sdk/execution_time/comms.py
+++ b/task_sdk/src/airflow/sdk/execution_time/comms.py
@@ -60,6 +60,7 @@ from airflow.sdk.api.datamodels._generated import (
     TIDeferredStatePayload,
     TIRescheduleStatePayload,
     TIRunContext,
+    TISuccessStatePayload,
     VariableResponse,
     XComResponse,
 )
@@ -191,11 +192,22 @@ class TaskState(BaseModel):
     - anything else = FAILED
     """
 
-    state: TerminalTIState
+    state: Literal[
+        TerminalTIState.FAILED,
+        TerminalTIState.SKIPPED,
+        TerminalTIState.REMOVED,
+        TerminalTIState.FAIL_WITHOUT_RETRY,
+    ]
     end_date: datetime | None = None
     type: Literal["TaskState"] = "TaskState"
 
 
+class SucceedTask(TISuccessStatePayload):
+    """Update a task's state to success. Includes task_outlets and 
outlet_events for registering asset events."""
+
+    type: Literal["SucceedTask"] = "SucceedTask"
+
+
 class DeferTask(TIDeferredStatePayload):
     """Update a task instance state to deferred."""
 
@@ -292,6 +304,7 @@ class GetPrevSuccessfulDagRun(BaseModel):
 
 ToSupervisor = Annotated[
     Union[
+        SucceedTask,
         DeferTask,
         GetAssetByName,
         GetAssetByUri,
diff --git a/task_sdk/src/airflow/sdk/execution_time/supervisor.py 
b/task_sdk/src/airflow/sdk/execution_time/supervisor.py
index 45da306722f..569855016cf 100644
--- a/task_sdk/src/airflow/sdk/execution_time/supervisor.py
+++ b/task_sdk/src/airflow/sdk/execution_time/supervisor.py
@@ -76,6 +76,7 @@ from airflow.sdk.execution_time.comms import (
     SetRenderedFields,
     SetXCom,
     StartupDetails,
+    SucceedTask,
     TaskState,
     ToSupervisor,
     VariableResult,
@@ -104,7 +105,11 @@ MAX_FAILED_HEARTBEATS: int = 3
 # These are the task instance states that require some additional information 
to transition into.
 # "Directly" here means that the PATCH API calls to transition into these 
states are
 # made from _handle_request() itself and don't have to come all the way to 
wait().
-STATES_SENT_DIRECTLY = [IntermediateTIState.DEFERRED, 
IntermediateTIState.UP_FOR_RESCHEDULE]
+STATES_SENT_DIRECTLY = [
+    IntermediateTIState.DEFERRED,
+    IntermediateTIState.UP_FOR_RESCHEDULE,
+    TerminalTIState.SUCCESS,
+]
 
 
 @overload
@@ -762,6 +767,14 @@ class ActivitySubprocess(WatchedSubprocess):
         if isinstance(msg, TaskState):
             self._terminal_state = msg.state
             self._task_end_time_monotonic = time.monotonic()
+        elif isinstance(msg, SucceedTask):
+            self._terminal_state = msg.state
+            self.client.task_instances.succeed(
+                id=self.id,
+                when=msg.end_date,
+                task_outlets=msg.task_outlets,
+                outlet_events=msg.outlet_events,
+            )
         elif isinstance(msg, GetConnection):
             conn = self.client.connections.get(msg.conn_id)
             if isinstance(conn, ConnectionResponse):
diff --git a/task_sdk/src/airflow/sdk/execution_time/task_runner.py 
b/task_sdk/src/airflow/sdk/execution_time/task_runner.py
index c35a79e13c3..c2d2c51b630 100644
--- a/task_sdk/src/airflow/sdk/execution_time/task_runner.py
+++ b/task_sdk/src/airflow/sdk/execution_time/task_runner.py
@@ -33,8 +33,9 @@ import structlog
 from pydantic import BaseModel, ConfigDict, Field, JsonValue, TypeAdapter
 
 from airflow.dag_processing.bundles.manager import DagBundlesManager
-from airflow.sdk.api.datamodels._generated import TaskInstance, 
TerminalTIState, TIRunContext
+from airflow.sdk.api.datamodels._generated import AssetProfile, TaskInstance, 
TerminalTIState, TIRunContext
 from airflow.sdk.definitions._internal.dag_parsing_context import 
_airflow_parsing_context_manager
+from airflow.sdk.definitions.asset import Asset, AssetAlias, AssetNameRef, 
AssetUriRef
 from airflow.sdk.definitions.baseoperator import BaseOperator
 from airflow.sdk.execution_time.comms import (
     DeferTask,
@@ -43,6 +44,7 @@ from airflow.sdk.execution_time.comms import (
     SetRenderedFields,
     SetXCom,
     StartupDetails,
+    SucceedTask,
     TaskState,
     ToSupervisor,
     ToTask,
@@ -446,6 +448,36 @@ def _get_rendered_fields(task: BaseOperator) -> dict[str, 
JsonValue]:
     return {field: serialize_template_field(getattr(task, field), field) for 
field in task.template_fields}
 
 
+def _process_outlets(context: Context, outlets: list[AssetProfile]):
+    added_alias_to_task_outlet = False
+    task_outlets: list[AssetProfile] = []
+    outlet_events: list[Any] = []
+    events = context["outlet_events"]
+
+    for obj in outlets or []:
+        # Lineage can have other types of objects besides assets
+        asset_type = type(obj).__name__
+        if isinstance(obj, Asset):
+            task_outlets.append(AssetProfile(name=obj.name, uri=obj.uri, 
asset_type=asset_type))
+            outlet_events.append(attrs.asdict(events[obj]))  # type: ignore
+        elif isinstance(obj, AssetNameRef):
+            task_outlets.append(AssetProfile(name=obj.name, 
asset_type=asset_type))
+            # Send all events, filtering can be done in API server.
+            outlet_events.append(attrs.asdict(events))  # type: ignore
+        elif isinstance(obj, AssetUriRef):
+            task_outlets.append(AssetProfile(uri=obj.uri, 
asset_type=asset_type))
+            # Send all events, filtering can be done in API server.
+            outlet_events.append(attrs.asdict(events))  # type: ignore
+        elif isinstance(obj, AssetAlias):
+            if not added_alias_to_task_outlet:
+                task_outlets.append(AssetProfile(asset_type=asset_type))
+                added_alias_to_task_outlet = True
+            for asset_alias_event in events[obj].asset_alias_events:
+                outlet_events.append(attrs.asdict(asset_alias_event))
+
+    return task_outlets, outlet_events
+
+
 def run(ti: RuntimeTaskInstance, log: Logger):
     """Run the task in this process."""
     from airflow.exceptions import (
@@ -477,12 +509,18 @@ def run(ti: RuntimeTaskInstance, log: Logger):
 
         _push_xcom_if_needed(result, ti)
 
+        task_outlets, outlet_events = _process_outlets(context, 
ti.task.outlets)
+
         # TODO: Get things from _execute_task_with_callbacks
         #   - Clearing XCom
         #   - Update RTIF
         #   - Pre Execute
         #   etc
-        msg = TaskState(state=TerminalTIState.SUCCESS, 
end_date=datetime.now(tz=timezone.utc))
+        msg = SucceedTask(
+            end_date=datetime.now(tz=timezone.utc),
+            task_outlets=task_outlets,
+            outlet_events=outlet_events,
+        )
     except TaskDeferred as defer:
         # TODO: Should we use structlog.bind_contextvars here for dag_id, 
task_id & run_id?
         log.info("Pausing task as DEFERRED. ", dag_id=ti.dag_id, 
task_id=ti.task_id, run_id=ti.run_id)
diff --git a/task_sdk/tests/execution_time/test_supervisor.py 
b/task_sdk/tests/execution_time/test_supervisor.py
index f2fcca8a2ab..ed921ed216a 100644
--- a/task_sdk/tests/execution_time/test_supervisor.py
+++ b/task_sdk/tests/execution_time/test_supervisor.py
@@ -55,6 +55,7 @@ from airflow.sdk.execution_time.comms import (
     RescheduleTask,
     SetRenderedFields,
     SetXCom,
+    SucceedTask,
     TaskState,
     VariableResult,
     XComResult,
@@ -978,6 +979,20 @@ class TestHandleRequest:
                 AssetResult(name="asset", uri="s3://bucket/obj", 
group="asset"),
                 id="get_asset_by_uri",
             ),
+            pytest.param(
+                SucceedTask(end_date=timezone.parse("2024-10-31T12:00:00Z")),
+                b"",
+                "task_instances.succeed",
+                (),
+                {
+                    "id": TI_ID,
+                    "outlet_events": None,
+                    "task_outlets": None,
+                    "when": timezone.parse("2024-10-31T12:00:00Z"),
+                },
+                "",
+                id="succeed_task",
+            ),
             pytest.param(
                 GetPrevSuccessfulDagRun(ti_id=TI_ID),
                 (
@@ -1002,6 +1017,7 @@ class TestHandleRequest:
         self,
         watched_subprocess,
         mocker,
+        time_machine,
         message,
         expected_buffer,
         client_attr_path,
@@ -1030,6 +1046,7 @@ class TestHandleRequest:
         next(generator)
         msg = message.model_dump_json().encode() + b"\n"
         generator.send(msg)
+        time_machine.move_to(timezone.datetime(2024, 10, 31), tick=False)
 
         # Verify the correct client method was called
         if client_attr_path:
diff --git a/task_sdk/tests/execution_time/test_task_runner.py 
b/task_sdk/tests/execution_time/test_task_runner.py
index 17c0029844e..e95d8cdf36b 100644
--- a/task_sdk/tests/execution_time/test_task_runner.py
+++ b/task_sdk/tests/execution_time/test_task_runner.py
@@ -37,7 +37,8 @@ from airflow.exceptions import (
     AirflowTaskTerminated,
 )
 from airflow.sdk import DAG, BaseOperator, Connection, get_current_context
-from airflow.sdk.api.datamodels._generated import TaskInstance, TerminalTIState
+from airflow.sdk.api.datamodels._generated import AssetProfile, TaskInstance, 
TerminalTIState
+from airflow.sdk.definitions.asset import Asset, AssetAlias
 from airflow.sdk.definitions.variable import Variable
 from airflow.sdk.execution_time.comms import (
     BundleInfo,
@@ -49,6 +50,7 @@ from airflow.sdk.execution_time.comms import (
     PrevSuccessfulDagRunResult,
     SetRenderedFields,
     StartupDetails,
+    SucceedTask,
     TaskState,
     VariableResult,
     XComResult,
@@ -172,7 +174,8 @@ def test_run_basic(time_machine, create_runtime_ti, 
spy_agency, mock_supervisor_
     assert ti.task._lock_for_execution
 
     mock_supervisor_comms.send_request.assert_called_once_with(
-        msg=TaskState(state=TerminalTIState.SUCCESS, end_date=instant), 
log=mock.ANY
+        msg=SucceedTask(state=TerminalTIState.SUCCESS, end_date=instant, 
task_outlets=[], outlet_events=[]),
+        log=mock.ANY,
     )
 
 
@@ -443,7 +446,12 @@ def test_startup_and_run_dag_with_rtif(
             log=mock.ANY,
         ),
         mock.call.send_request(
-            msg=TaskState(end_date=instant, state=TerminalTIState.SUCCESS),
+            msg=SucceedTask(
+                end_date=instant,
+                state=TerminalTIState.SUCCESS,
+                task_outlets=[],
+                outlet_events=[],
+            ),
             log=mock.ANY,
         ),
     ]
@@ -498,7 +506,8 @@ def test_get_context_in_task(create_runtime_ti, 
time_machine, mock_supervisor_co
 
     # Ensure the task is Successful
     mock_supervisor_comms.send_request.assert_called_once_with(
-        msg=TaskState(state=TerminalTIState.SUCCESS, end_date=instant), 
log=mock.ANY
+        msg=SucceedTask(state=TerminalTIState.SUCCESS, end_date=instant, 
task_outlets=[], outlet_events=[]),
+        log=mock.ANY,
     )
 
 
@@ -588,6 +597,60 @@ def test_dag_parsing_context(make_ti_context, 
mock_supervisor_comms, monkeypatch
     assert ti.task.dag.task_dict.keys() == {"visible_task", "conditional_task"}
 
 
[email protected](
+    ["task_outlets", "expected_msg"],
+    [
+        pytest.param(
+            [Asset(name="s3://bucket/my-task", uri="s3://bucket/my-task")],
+            SucceedTask(
+                state="success",
+                end_date=timezone.datetime(2024, 12, 3, 10, 0),
+                task_outlets=[
+                    AssetProfile(name="s3://bucket/my-task", 
uri="s3://bucket/my-task", asset_type="Asset")
+                ],
+                outlet_events=[
+                    {
+                        "key": {"name": "s3://bucket/my-task", "uri": 
"s3://bucket/my-task"},
+                        "extra": {},
+                        "asset_alias_events": [],
+                    }
+                ],
+            ),
+            id="asset",
+        ),
+        pytest.param(
+            [AssetAlias(name="example-alias", group="asset")],
+            SucceedTask(
+                state="success",
+                end_date=timezone.datetime(2024, 12, 3, 10, 0),
+                task_outlets=[AssetProfile(asset_type="AssetAlias")],
+                outlet_events=[],
+            ),
+            id="asset-alias",
+        ),
+    ],
+)
+def test_run_with_asset_outlets(
+    time_machine, create_runtime_ti, mock_supervisor_comms, task_outlets, 
expected_msg
+):
+    """Test running a basic task that contains asset outlets."""
+    from airflow.providers.standard.operators.bash import BashOperator
+
+    task = BashOperator(
+        outlets=task_outlets,
+        task_id="asset-outlet-task",
+        bash_command="echo 'hi'",
+    )
+
+    ti = create_runtime_ti(task=task, dag_id="dag_with_asset_outlet_task")
+    instant = timezone.datetime(2024, 12, 3, 10, 0)
+    time_machine.move_to(instant, tick=False)
+
+    run(ti, log=mock.MagicMock())
+
+    mock_supervisor_comms.send_request.assert_any_call(msg=expected_msg, 
log=mock.ANY)
+
+
 class TestRuntimeTaskInstance:
     def test_get_context_without_ti_context_from_server(self, mocked_parse, 
make_ti_context):
         """Test get_template_context without ti_context_from_server."""
diff --git a/tests/api_fastapi/execution_api/routes/test_task_instances.py 
b/tests/api_fastapi/execution_api/routes/test_task_instances.py
index e3aef1505bc..9ccd1b0d088 100644
--- a/tests/api_fastapi/execution_api/routes/test_task_instances.py
+++ b/tests/api_fastapi/execution_api/routes/test_task_instances.py
@@ -26,11 +26,12 @@ from sqlalchemy import select, update
 from sqlalchemy.exc import SQLAlchemyError
 
 from airflow.models import RenderedTaskInstanceFields, TaskReschedule, Trigger
+from airflow.models.asset import AssetActive, AssetAliasModel, AssetEvent, 
AssetModel
 from airflow.models.taskinstance import TaskInstance
 from airflow.utils import timezone
 from airflow.utils.state import State, TaskInstanceState, TerminalTIState
 
-from tests_common.test_utils.db import clear_db_runs, clear_rendered_ti_fields
+from tests_common.test_utils.db import clear_db_assets, clear_db_runs, 
clear_rendered_ti_fields
 
 pytestmark = pytest.mark.db_test
 
@@ -39,6 +40,19 @@ DEFAULT_START_DATE = timezone.parse("2024-10-31T11:00:00Z")
 DEFAULT_END_DATE = timezone.parse("2024-10-31T12:00:00Z")
 
 
+def _create_asset_aliases(session, num: int = 2) -> None:
+    asset_aliases = [
+        AssetAliasModel(
+            id=i,
+            name=f"simple{i}",
+            group="alias",
+        )
+        for i in range(1, 1 + num)
+    ]
+    session.add_all(asset_aliases)
+    session.commit()
+
+
 class TestTIRunState:
     def setup_method(self):
         clear_db_runs()
@@ -267,6 +281,87 @@ class TestTIUpdateState:
         assert ti.state == expected_state
         assert ti.end_date == end_date
 
+    @pytest.mark.parametrize(
+        ("task_outlets", "outlet_events"),
+        [
+            (
+                [{"name": "s3://bucket/my-task", "uri": "s3://bucket/my-task", 
"asset_type": "Asset"}],
+                [
+                    {
+                        "key": {"name": "s3://bucket/my-task", "uri": 
"s3://bucket/my-task"},
+                        "extra": {},
+                        "asset_alias_events": [],
+                    }
+                ],
+            ),
+            (
+                [{"asset_type": "AssetAlias"}],
+                [
+                    {
+                        "source_alias_name": "example-alias",
+                        "dest_asset_key": {"name": "s3://bucket/my-task", 
"uri": "s3://bucket/my-task"},
+                        "extra": {},
+                    }
+                ],
+            ),
+        ],
+    )
+    def test_ti_update_state_to_success_with_asset_events(
+        self, client, session, create_task_instance, task_outlets, 
outlet_events
+    ):
+        clear_db_assets()
+        clear_db_runs()
+
+        asset = AssetModel(
+            id=1,
+            name="s3://bucket/my-task",
+            uri="s3://bucket/my-task",
+            group="asset",
+            extra={},
+        )
+        asset_active = AssetActive.for_asset(asset)
+        session.add_all([asset, asset_active])
+        asset_type = task_outlets[0]["asset_type"]
+        if asset_type == "AssetAlias":
+            _create_asset_aliases(session, num=1)
+            asset_alias = session.query(AssetAliasModel).all()
+            assert len(asset_alias) == 1
+            assert asset_alias == [AssetAliasModel(name="simple1")]
+
+        ti = create_task_instance(
+            task_id="test_ti_update_state_to_success_with_asset_events",
+            start_date=DEFAULT_START_DATE,
+            state=State.RUNNING,
+        )
+        session.commit()
+
+        response = client.patch(
+            f"/execution/task-instances/{ti.id}/state",
+            json={
+                "state": "success",
+                "end_date": DEFAULT_END_DATE.isoformat(),
+                "task_outlets": task_outlets,
+                "outlet_events": outlet_events,
+            },
+        )
+
+        assert response.status_code == 204
+        assert response.text == ""
+        session.expire_all()
+
+        # check if asset was created properly
+        asset = session.query(AssetModel).all()
+        assert len(asset) == 1
+        assert asset == [AssetModel(name="s3://bucket/my-task", 
uri="s3://bucket/my-task", extra={})]
+
+        event = session.query(AssetEvent).all()
+        assert len(event) == 1
+        assert event[0].asset_id == 1
+        assert event[0].asset == AssetModel(name="s3://bucket/my-task", 
uri="s3://bucket/my-task", extra={})
+        assert event[0].extra == {}
+        if asset_type == "AssetAlias":
+            assert event[0].source_aliases == 
[AssetAliasModel(name="example-alias")]
+
     def test_ti_update_state_not_found(self, client, session):
         """
         Test that a 404 error is returned when the Task Instance does not 
exist.
@@ -319,13 +414,20 @@ class TestTIUpdateState:
             "end_date": "2024-10-31T12:00:00Z",
         }
 
-        with mock.patch(
-            "airflow.api_fastapi.common.db.common.Session.execute",
-            side_effect=[
-                mock.Mock(one=lambda: ("running", 1, 0)),  # First call 
returns "queued"
-                SQLAlchemyError("Database error"),  # Second call raises an 
error
-            ],
+        with (
+            mock.patch(
+                "airflow.api_fastapi.common.db.common.Session.execute",
+                side_effect=[
+                    mock.Mock(one=lambda: ("running", 1, 0)),  # First call 
returns "queued"
+                    mock.Mock(one=lambda: ("running", 1, 0)),  # Second call 
returns "queued"
+                    SQLAlchemyError("Database error"),  # Last call raises an 
error
+                ],
+            ),
+            mock.patch(
+                
"airflow.models.taskinstance.TaskInstance.register_asset_changes_in_db",
+            ) as mock_register_asset_changes_in_db,
         ):
+            mock_register_asset_changes_in_db.return_value = None
             response = 
client.patch(f"/execution/task-instances/{ti.id}/state", json=payload)
             assert response.status_code == 500
             assert response.json()["detail"] == "Database error occurred"


Reply via email to