|
14 | 14 |
|
15 | 15 | from airflow import configuration |
16 | 16 |
|
| 17 | +from dagfactory.utils import check_dict_key |
| 18 | + |
17 | 19 | try: |
18 | 20 | from airflow.sdk.bases.operator import BaseOperator |
19 | 21 | from airflow.sdk.definitions.dag import DAG |
@@ -150,6 +152,18 @@ def get_dag_params(self) -> Dict[str, Any]: |
150 | 152 | dag_params["dagrun_timeout"]: timedelta = timedelta(seconds=dag_params["dagrun_timeout_sec"]) |
151 | 153 | del dag_params["dagrun_timeout_sec"] |
152 | 154 |
|
| 155 | + if utils.check_dict_key(dag_params, "start_date"): |
| 156 | + dag_params["start_date"]: datetime = utils.get_datetime( |
| 157 | + date_value=dag_params["start_date"], |
| 158 | + timezone=dag_params.get("timezone", "UTC"), |
| 159 | + ) |
| 160 | + |
| 161 | + if utils.check_dict_key(dag_params, "end_date"): |
| 162 | + dag_params["end_date"]: datetime = utils.get_datetime( |
| 163 | + date_value=dag_params["end_date"], |
| 164 | + timezone=dag_params.get("timezone", "UTC"), |
| 165 | + ) |
| 166 | + |
153 | 167 | # Convert from 'end_date: Union[str, datetime, date]' to 'end_date: datetime' |
154 | 168 | if utils.check_dict_key(dag_params["default_args"], "end_date"): |
155 | 169 | dag_params["default_args"]["end_date"]: datetime = utils.get_datetime( |
@@ -235,10 +249,11 @@ def get_dag_params(self) -> Dict[str, Any]: |
235 | 249 | try: |
236 | 250 | # ensure that default_args dictionary contains key "start_date" |
237 | 251 | # 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 | | - ) |
| 252 | + if check_dict_key(dag_params["default_args"], "start_date"): |
| 253 | + dag_params["default_args"]["start_date"]: datetime = utils.get_datetime( |
| 254 | + date_value=dag_params["default_args"]["start_date"], |
| 255 | + timezone=dag_params["default_args"].get("timezone", "UTC"), |
| 256 | + ) |
242 | 257 | except KeyError as err: |
243 | 258 | # pylint: disable=line-too-long |
244 | 259 | raise DagFactoryConfigException(f"{self.dag_name} config is missing start_date") from err |
@@ -985,6 +1000,9 @@ def build(self) -> Dict[str, Union[str, DAG]]: |
985 | 1000 |
|
986 | 1001 | dag_kwargs["params"] = dag_params.get("params", None) |
987 | 1002 |
|
| 1003 | + dag_kwargs["start_date"] = dag_params.get("start_date", None) |
| 1004 | + dag_kwargs["end_date"] = dag_params.get("end_date", None) |
| 1005 | + |
988 | 1006 | dag: DAG = DAG(**dag_kwargs) |
989 | 1007 |
|
990 | 1008 | if dag_params.get("doc_md_file_path"): |
|
0 commit comments