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"