Skip to content

Commit ddb5246

Browse files
authored
Refactor operator links to not create ad hoc TaskInstances (#21285)
1 parent dc3c47d commit ddb5246

7 files changed

Lines changed: 33 additions & 27 deletions

File tree

airflow/models/xcom.py

Lines changed: 4 additions & 4 deletions
Original file line numberDiff line numberDiff line change
@@ -220,7 +220,7 @@ def get_one(
220220
@classmethod
221221
def get_one(
222222
cls,
223-
execution_date: pendulum.DateTime,
223+
execution_date: datetime.datetime,
224224
key: Optional[str] = None,
225225
task_id: Optional[str] = None,
226226
dag_id: Optional[str] = None,
@@ -233,7 +233,7 @@ def get_one(
233233
@provide_session
234234
def get_one(
235235
cls,
236-
execution_date: Optional[pendulum.DateTime] = None,
236+
execution_date: Optional[datetime.datetime] = None,
237237
key: Optional[str] = None,
238238
task_id: Optional[Union[str, Iterable[str]]] = None,
239239
dag_id: Optional[Union[str, Iterable[str]]] = None,
@@ -314,7 +314,7 @@ def get_many(
314314
@classmethod
315315
def get_many(
316316
cls,
317-
execution_date: pendulum.DateTime,
317+
execution_date: datetime.datetime,
318318
key: Optional[str] = None,
319319
task_ids: Union[str, Iterable[str], None] = None,
320320
dag_ids: Union[str, Iterable[str], None] = None,
@@ -328,7 +328,7 @@ def get_many(
328328
@provide_session
329329
def get_many(
330330
cls,
331-
execution_date: Optional[pendulum.DateTime] = None,
331+
execution_date: Optional[datetime.datetime] = None,
332332
key: Optional[str] = None,
333333
task_ids: Optional[Union[str, Iterable[str]]] = None,
334334
dag_ids: Optional[Union[str, Iterable[str]]] = None,

airflow/providers/amazon/aws/operators/emr.py

Lines changed: 4 additions & 3 deletions
Original file line numberDiff line numberDiff line change
@@ -22,7 +22,7 @@
2222
from uuid import uuid4
2323

2424
from airflow.exceptions import AirflowException
25-
from airflow.models import BaseOperator, BaseOperatorLink, TaskInstance
25+
from airflow.models import BaseOperator, BaseOperatorLink, XCom
2626
from airflow.providers.amazon.aws.hooks.emr import EmrHook
2727

2828
if TYPE_CHECKING:
@@ -238,8 +238,9 @@ def get_link(self, operator: BaseOperator, dttm: datetime) -> str:
238238
:param dttm: datetime
239239
:return: url link
240240
"""
241-
ti = TaskInstance(task=operator, execution_date=dttm)
242-
flow_id = ti.xcom_pull(task_ids=operator.task_id)
241+
flow_id = XCom.get_one(
242+
key="return_value", dag_id=operator.dag.dag_id, task_id=operator.task_id, execution_date=dttm
243+
)
243244
return (
244245
f'https://console.aws.amazon.com/elasticmapreduce/home#cluster-details:{flow_id}'
245246
if flow_id

airflow/providers/google/cloud/operators/bigquery.py

Lines changed: 3 additions & 3 deletions
Original file line numberDiff line numberDiff line change
@@ -32,7 +32,6 @@
3232

3333
from airflow.exceptions import AirflowException
3434
from airflow.models import BaseOperator, BaseOperatorLink
35-
from airflow.models.taskinstance import TaskInstance
3635
from airflow.models.xcom import XCom
3736
from airflow.operators.sql import SQLCheckOperator, SQLIntervalCheckOperator, SQLValueCheckOperator
3837
from airflow.providers.google.cloud.hooks.bigquery import BigQueryHook, BigQueryJob
@@ -84,8 +83,9 @@ def name(self) -> str:
8483
return f'BigQuery Console #{self.index + 1}'
8584

8685
def get_link(self, operator: BaseOperator, dttm: datetime):
87-
ti = TaskInstance(task=operator, execution_date=dttm)
88-
job_ids = ti.xcom_pull(task_ids=operator.task_id, key='job_id')
86+
job_ids = XCom.get_one(
87+
key='job_id', dag_id=operator.dag.dag_id, task_id=operator.task_id, execution_date=dttm
88+
)
8989
if not job_ids:
9090
return None
9191
if len(job_ids) < self.index:

airflow/providers/google/cloud/operators/dataproc.py

Lines changed: 7 additions & 6 deletions
Original file line numberDiff line numberDiff line change
@@ -36,8 +36,7 @@
3636
from google.protobuf.field_mask_pb2 import FieldMask
3737

3838
from airflow.exceptions import AirflowException
39-
from airflow.models import BaseOperator, BaseOperatorLink
40-
from airflow.models.taskinstance import TaskInstance
39+
from airflow.models import BaseOperator, BaseOperatorLink, XCom
4140
from airflow.providers.google.cloud.hooks.dataproc import DataprocHook, DataProcJobBuilder
4241
from airflow.providers.google.cloud.hooks.gcs import GCSHook
4342
from airflow.utils import timezone
@@ -59,8 +58,9 @@ class DataprocJobLink(BaseOperatorLink):
5958
name = "Dataproc Job"
6059

6160
def get_link(self, operator, dttm):
62-
ti = TaskInstance(task=operator, execution_date=dttm)
63-
job_conf = ti.xcom_pull(task_ids=operator.task_id, key="job_conf")
61+
job_conf = XCom.get_one(
62+
key="job_conf", dag_id=operator.dag.dag_id, task_id=operator.task_id, execution_date=dttm
63+
)
6464
return (
6565
DATAPROC_JOB_LOG_LINK.format(
6666
job_id=job_conf["job_id"],
@@ -78,8 +78,9 @@ class DataprocClusterLink(BaseOperatorLink):
7878
name = "Dataproc Cluster"
7979

8080
def get_link(self, operator, dttm):
81-
ti = TaskInstance(task=operator, execution_date=dttm)
82-
cluster_conf = ti.xcom_pull(task_ids=operator.task_id, key="cluster_conf")
81+
cluster_conf = XCom.get_one(
82+
key="cluster_conf", dag_id=operator.dag.dag_id, task_id=operator.task_id, execution_date=dttm
83+
)
8384
return (
8485
DATAPROC_CLUSTER_LINK.format(
8586
cluster_name=cluster_conf["cluster_name"],

airflow/providers/google/cloud/operators/mlengine.py

Lines changed: 4 additions & 4 deletions
Original file line numberDiff line numberDiff line change
@@ -22,8 +22,7 @@
2222
from typing import TYPE_CHECKING, Any, Dict, List, Optional, Sequence, Union
2323

2424
from airflow.exceptions import AirflowException
25-
from airflow.models import BaseOperator, BaseOperatorLink
26-
from airflow.models.taskinstance import TaskInstance
25+
from airflow.models import BaseOperator, BaseOperatorLink, XCom
2726
from airflow.providers.google.cloud.hooks.mlengine import MLEngineHook
2827

2928
if TYPE_CHECKING:
@@ -980,8 +979,9 @@ class AIPlatformConsoleLink(BaseOperatorLink):
980979
name = "AI Platform Console"
981980

982981
def get_link(self, operator, dttm):
983-
task_instance = TaskInstance(task=operator, execution_date=dttm)
984-
gcp_metadata_dict = task_instance.xcom_pull(task_ids=operator.task_id, key="gcp_metadata")
982+
gcp_metadata_dict = XCom.get_one(
983+
key="gcp_metadata", dag_id=operator.dag.dag_id, task_id=operator.task_id, execution_date=dttm
984+
)
985985
if not gcp_metadata_dict:
986986
return ''
987987
job_id = gcp_metadata_dict['job_id']

airflow/providers/microsoft/azure/operators/data_factory.py

Lines changed: 7 additions & 3 deletions
Original file line numberDiff line numberDiff line change
@@ -18,7 +18,7 @@
1818
from typing import TYPE_CHECKING, Any, Dict, Optional, Sequence
1919

2020
from airflow.hooks.base import BaseHook
21-
from airflow.models import BaseOperator, BaseOperatorLink, TaskInstance
21+
from airflow.models import BaseOperator, BaseOperatorLink, XCom
2222
from airflow.providers.microsoft.azure.hooks.data_factory import (
2323
AzureDataFactoryHook,
2424
AzureDataFactoryPipelineRunException,
@@ -35,8 +35,12 @@ class AzureDataFactoryPipelineRunLink(BaseOperatorLink):
3535
name = "Monitor Pipeline Run"
3636

3737
def get_link(self, operator, dttm):
38-
ti = TaskInstance(task=operator, execution_date=dttm)
39-
run_id = ti.xcom_pull(task_ids=operator.task_id, key="run_id")
38+
run_id = XCom.get_one(
39+
key="run_id",
40+
dag_id=operator.dag.dag_id,
41+
task_id=operator.task_id,
42+
execution_date=dttm,
43+
)
4044

4145
conn = BaseHook.get_connection(operator.azure_data_factory_conn_id)
4246
subscription_id = conn.extra_dejson["extra__azure_data_factory__subscriptionId"]

airflow/providers/qubole/operators/qubole.py

Lines changed: 4 additions & 4 deletions
Original file line numberDiff line numberDiff line change
@@ -21,8 +21,7 @@
2121
from typing import TYPE_CHECKING, Optional, Sequence
2222

2323
from airflow.hooks.base import BaseHook
24-
from airflow.models import BaseOperator, BaseOperatorLink
25-
from airflow.models.taskinstance import TaskInstance
24+
from airflow.models import BaseOperator, BaseOperatorLink, XCom
2625
from airflow.providers.qubole.hooks.qubole import (
2726
COMMAND_ARGS,
2827
HYPHEN_ARGS,
@@ -48,7 +47,6 @@ def get_link(self, operator: BaseOperator, dttm: datetime) -> str:
4847
:param dttm: datetime
4948
:return: url link
5049
"""
51-
ti = TaskInstance(task=operator, execution_date=dttm)
5250
conn = BaseHook.get_connection(
5351
getattr(operator, "qubole_conn_id", None)
5452
or operator.kwargs['qubole_conn_id'] # type: ignore[attr-defined]
@@ -57,7 +55,9 @@ def get_link(self, operator: BaseOperator, dttm: datetime) -> str:
5755
host = re.sub(r'api$', 'v2/analyze?command_id=', conn.host)
5856
else:
5957
host = 'https://api.qubole.com/v2/analyze?command_id='
60-
qds_command_id = ti.xcom_pull(task_ids=operator.task_id, key='qbol_cmd_id')
58+
qds_command_id = XCom.get_one(
59+
key='qbol_cmd_id', dag_id=operator.dag.dag_id, task_id=operator.task_id, execution_date=dttm
60+
)
6161
url = host + str(qds_command_id) if qds_command_id else ''
6262
return url
6363

0 commit comments

Comments
 (0)