|
8 | 8 | import re |
9 | 9 | import warnings |
10 | 10 | from copy import deepcopy |
11 | | -from datetime import datetime, timedelta |
| 11 | +from datetime import timedelta |
12 | 12 | from functools import partial, reduce |
13 | 13 | from typing import Any, Callable, Dict, List, Tuple, Union |
14 | 14 |
|
15 | 15 | from airflow import configuration |
16 | 16 |
|
| 17 | +from dagfactory.utils import convert_to_datetime_datetime, convert_to_timedelta |
| 18 | + |
17 | 19 | try: |
18 | 20 | from airflow.sdk.bases.operator import BaseOperator |
19 | 21 | from airflow.sdk.definitions.dag import DAG |
@@ -145,27 +147,28 @@ def get_dag_params(self) -> Dict[str, Any]: |
145 | 147 | if utils.check_dict_key(dag_params, "schedule_interval") and dag_params["schedule_interval"] == "None": |
146 | 148 | dag_params["schedule_interval"] = None |
147 | 149 |
|
148 | | - # Convert from 'dagrun_timeout_sec: int' to 'dagrun_timeout: timedelta' |
149 | | - if utils.check_dict_key(dag_params, "dagrun_timeout_sec"): |
150 | | - dag_params["dagrun_timeout"]: timedelta = timedelta(seconds=dag_params["dagrun_timeout_sec"]) |
151 | | - del dag_params["dagrun_timeout_sec"] |
| 150 | + # Convert from 'dagrun_timeout: int' to 'dagrun_timeout: timedelta' |
| 151 | + if utils.check_dict_key(dag_params, "dagrun_timeout"): |
| 152 | + if isinstance(dag_params["dagrun_timeout"], int): |
| 153 | + dag_params["dagrun_timeout"]: timedelta = timedelta(seconds=dag_params["dagrun_timeout"]) |
152 | 154 |
|
153 | | - # Convert from 'end_date: Union[str, datetime, date]' to 'end_date: datetime' |
154 | | - if utils.check_dict_key(dag_params["default_args"], "end_date"): |
155 | | - dag_params["default_args"]["end_date"]: datetime = utils.get_datetime( |
156 | | - date_value=dag_params["default_args"]["end_date"], |
157 | | - timezone=dag_params["default_args"].get("timezone", "UTC"), |
| 155 | + if dag_params["default_args"]["start_date"] and isinstance(dag_params["default_args"]["start_date"], str): |
| 156 | + dag_params["default_args"]["start_date"] = convert_to_datetime_datetime( |
| 157 | + dag_params["default_args"]["start_date"] |
158 | 158 | ) |
159 | 159 |
|
160 | | - if utils.check_dict_key(dag_params["default_args"], "retry_delay_sec"): |
161 | | - dag_params["default_args"]["retry_delay"]: timedelta = timedelta( |
162 | | - seconds=dag_params["default_args"]["retry_delay_sec"] |
| 160 | + if dag_params["default_args"]["end_date"] and isinstance(dag_params["default_args"]["end_date"], str): |
| 161 | + dag_params["default_args"]["end_date"] = convert_to_datetime_datetime( |
| 162 | + dag_params["default_args"]["end_date"] |
163 | 163 | ) |
164 | | - del dag_params["default_args"]["retry_delay_sec"] |
165 | 164 |
|
166 | | - if utils.check_dict_key(dag_params["default_args"], "sla_secs"): |
167 | | - dag_params["default_args"]["sla"]: timedelta = timedelta(seconds=dag_params["default_args"]["sla_secs"]) |
168 | | - del dag_params["default_args"]["sla_secs"] |
| 165 | + if utils.check_dict_key(dag_params["default_args"], "retry_delay"): |
| 166 | + if isinstance(dag_params["default_args"]["retry_delay"], int): |
| 167 | + dag_params["default_args"]["retry_delay"] = timedelta(seconds=dag_params["default_args"]["retry"]) |
| 168 | + |
| 169 | + if utils.check_dict_key(dag_params["default_args"], "sla"): |
| 170 | + if isinstance(dag_params["default_args"]["sla"], int): |
| 171 | + dag_params["default_args"]["sla"]: timedelta = timedelta(seconds=dag_params["default_args"]["sla"]) |
169 | 172 |
|
170 | 173 | # Parse callbacks at the DAG-level and at the Task-level, configured in default_args. Note that the version |
171 | 174 | # check has gone into the set_callback method |
@@ -232,16 +235,12 @@ def get_dag_params(self) -> Dict[str, Any]: |
232 | 235 | else: |
233 | 236 | raise DagFactoryException("render_template_as_native_obj should be bool type!") |
234 | 237 |
|
235 | | - try: |
236 | | - # ensure that default_args dictionary contains key "start_date" |
237 | | - # with "datetime" value in specified timezone |
238 | | - dag_params["default_args"]["start_date"]: datetime = utils.get_datetime( |
239 | | - date_value=dag_params["default_args"]["start_date"], |
240 | | - timezone=dag_params["default_args"].get("timezone", "UTC"), |
241 | | - ) |
242 | | - except KeyError as err: |
243 | | - # pylint: disable=line-too-long |
244 | | - raise DagFactoryConfigException(f"{self.dag_name} config is missing start_date") from err |
| 238 | + if dag_params.get("start_date") and isinstance(dag_params["start_date"], str): |
| 239 | + dag_params["start_date"] = convert_to_datetime_datetime(dag_params["start_date"]) |
| 240 | + |
| 241 | + if dag_params.get("end_date") and isinstance(dag_params["end_date"], str): |
| 242 | + dag_params["end_date"] = convert_to_datetime_datetime(dag_params["end_date"]) |
| 243 | + |
245 | 244 | return dag_params |
246 | 245 |
|
247 | 246 | @staticmethod |
@@ -854,10 +853,7 @@ def build(self) -> Dict[str, Union[str, DAG]]: |
854 | 853 | ) |
855 | 854 |
|
856 | 855 | if dag_params.get("timetable"): |
857 | | - timetable_args = dag_params.get("timetable") |
858 | | - dag_kwargs["timetable"] = DagBuilder.make_timetable( |
859 | | - timetable_args.get("callable"), timetable_args.get("params") |
860 | | - ) |
| 856 | + dag_kwargs["timetable"] = dag_params.get("timetable") |
861 | 857 |
|
862 | 858 | dag_kwargs["catchup"] = dag_params.get( |
863 | 859 | "catchup", configuration.conf.getboolean("scheduler", "catchup_by_default") |
@@ -1020,27 +1016,27 @@ def topological_sort_tasks(tasks_configs: dict[str, Any]) -> list[tuple(str, Any |
1020 | 1016 | return sorted_tasks |
1021 | 1017 |
|
1022 | 1018 | @staticmethod |
1023 | | - def adjust_general_task_params(task_params: dict(str, Any)): |
| 1019 | + def adjust_general_task_params(task_params: dict[str, Any]): |
1024 | 1020 | """Adjusts in place the task params argument""" |
1025 | | - if utils.check_dict_key(task_params, "execution_timeout_secs"): |
1026 | | - task_params["execution_timeout"]: timedelta = timedelta(seconds=task_params["execution_timeout_secs"]) |
1027 | | - del task_params["execution_timeout_secs"] |
| 1021 | + if utils.check_dict_key(task_params, "execution_timeout"): |
| 1022 | + if isinstance(task_params["execution_timeout"], int): |
| 1023 | + task_params["execution_timeout"]: timedelta = timedelta(seconds=task_params["execution_timeout"]) |
1028 | 1024 |
|
1029 | | - if utils.check_dict_key(task_params, "sla_secs"): |
1030 | | - task_params["sla"]: timedelta = timedelta(seconds=task_params["sla_secs"]) |
1031 | | - del task_params["sla_secs"] |
| 1025 | + if utils.check_dict_key(task_params, "sla"): |
| 1026 | + if isinstance(task_params["sla"], int): |
| 1027 | + task_params["sla"]: timedelta = timedelta(seconds=task_params["sla"]) |
1032 | 1028 |
|
1033 | | - if utils.check_dict_key(task_params, "execution_delta_secs"): |
1034 | | - task_params["execution_delta"]: timedelta = timedelta(seconds=task_params["execution_delta_secs"]) |
1035 | | - del task_params["execution_delta_secs"] |
| 1029 | + if utils.check_dict_key(task_params, "execution_delta"): |
| 1030 | + if isinstance(task_params["execution_delta"], int): |
| 1031 | + task_params["execution_delta"]: timedelta = timedelta(seconds=task_params["execution_delta"]) |
1036 | 1032 |
|
1037 | 1033 | # Used by airflow.sensors.external_task_sensor.ExternalTaskSensor |
1038 | 1034 | if utils.check_dict_key(task_params, "execution_date_fn"): |
1039 | 1035 | python_callable: Callable = import_string(task_params["execution_date_fn"]) |
1040 | 1036 | task_params["execution_date_fn"] = python_callable |
1041 | 1037 | elif utils.check_dict_key(task_params, "execution_delta"): |
1042 | | - execution_delta = utils.get_time_delta(task_params["execution_delta"]) |
1043 | | - task_params["execution_delta"] = execution_delta |
| 1038 | + task_params["execution_delta"] = convert_to_timedelta(task_params["execution_delta"]) |
| 1039 | + |
1044 | 1040 | elif utils.check_dict_key(task_params, "execution_date_fn_name") and utils.check_dict_key( |
1045 | 1041 | task_params, "execution_date_fn_file" |
1046 | 1042 | ): |
@@ -1174,7 +1170,7 @@ def make_decorator( |
1174 | 1170 | return decorator(**decorator_kwargs)(**callable_kwargs) |
1175 | 1171 |
|
1176 | 1172 | @staticmethod |
1177 | | - def replace_kwargs_values_as_tasks(kwargs: dict(str, Any), tasks_dict: dict(str, Any)): |
| 1173 | + def replace_kwargs_values_as_tasks(kwargs: dict[str, Any], tasks_dict: dict[str, Any]): |
1178 | 1174 | for key, value in kwargs.items(): |
1179 | 1175 | if isinstance(value, str) and value.startswith("+"): |
1180 | 1176 | upstream_task_name = value.split("+")[-1] |
|
0 commit comments