Skip to content

Commit 4babc47

Browse files
committed
Fix start_date, end_date dag level uses
1 parent 5ed0b48 commit 4babc47

2 files changed

Lines changed: 51 additions & 5 deletions

File tree

dagfactory/dagbuilder.py

Lines changed: 22 additions & 4 deletions
Original file line numberDiff line numberDiff line change
@@ -14,6 +14,8 @@
1414

1515
from airflow import configuration
1616

17+
from dagfactory.utils import check_dict_key
18+
1719
try:
1820
from airflow.sdk.bases.operator import BaseOperator
1921
from airflow.sdk.definitions.dag import DAG
@@ -150,6 +152,18 @@ def get_dag_params(self) -> Dict[str, Any]:
150152
dag_params["dagrun_timeout"]: timedelta = timedelta(seconds=dag_params["dagrun_timeout_sec"])
151153
del dag_params["dagrun_timeout_sec"]
152154

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+
153167
# Convert from 'end_date: Union[str, datetime, date]' to 'end_date: datetime'
154168
if utils.check_dict_key(dag_params["default_args"], "end_date"):
155169
dag_params["default_args"]["end_date"]: datetime = utils.get_datetime(
@@ -235,10 +249,11 @@ def get_dag_params(self) -> Dict[str, Any]:
235249
try:
236250
# ensure that default_args dictionary contains key "start_date"
237251
# 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+
)
242257
except KeyError as err:
243258
# pylint: disable=line-too-long
244259
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]]:
9851000

9861001
dag_kwargs["params"] = dag_params.get("params", None)
9871002

1003+
dag_kwargs["start_date"] = dag_params.get("start_date", None)
1004+
dag_kwargs["end_date"] = dag_params.get("end_date", None)
1005+
9881006
dag: DAG = DAG(**dag_kwargs)
9891007

9901008
if dag_params.get("doc_md_file_path"):

tests/test_dagfactory.py

Lines changed: 29 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -1,6 +1,7 @@
11
import datetime
22
import logging
33
import os
4+
import tempfile
45

56
import pytest
67
from airflow.version import version as AIRFLOW_VERSION
@@ -11,12 +12,13 @@
1112
from airflow.models.variable import Variable # noqa: F401
1213

1314
from packaging import version
15+
from pendulum.datetime import DateTime, Timezone
1416

1517
from tests.utils import get_bash_operator_path, get_schedule_key
1618

1719
here = os.path.dirname(__file__)
1820

19-
from dagfactory import dagfactory, load_yaml_dags
21+
from dagfactory import DagFactory, dagfactory, load_yaml_dags
2022

2123
TEST_DAG_FACTORY = os.path.join(here, "fixtures/dag_factory.yml")
2224
DAG_FACTORY_NO_OR_NONE_STRING_SCHEDULE = os.path.join(here, "fixtures/dag_factory_no_or_none_string_schedule.yml")
@@ -592,3 +594,29 @@ def test_yml_dag_rendering_in_docs():
592594
with open(dag_path, "r") as file:
593595
expected_doc_md = "## YML DAG\n```yaml\n" + file.read() + "\n```"
594596
assert generated_doc_md == expected_doc_md
597+
598+
599+
def test_dag_level_start():
600+
data = """
601+
my_dag:
602+
schedule: "0 3 * * *"
603+
start_date: 2024-11-11
604+
end_date: 2025-11-11
605+
tasks:
606+
task_1:
607+
operator: airflow.operators.bash.BashOperator
608+
bash_command: "echo 1"
609+
"""
610+
611+
# Write to temporary YAML file
612+
with tempfile.NamedTemporaryFile(mode="w", suffix=".yaml", delete=False) as tmp:
613+
tmp.write(data)
614+
temp_file = tmp.name
615+
616+
# Use DagFactory to load and generate DAGs
617+
df = DagFactory(config_filepath=temp_file)
618+
df.generate_dags(globals=globals())
619+
dag = globals()["my_dag"]
620+
621+
assert dag.start_date == DateTime(2024, 11, 11, 0, 0, 0, tzinfo=Timezone("UTC"))
622+
assert dag.end_date == DateTime(2025, 11, 11, 0, 0, 0, tzinfo=Timezone("UTC"))

0 commit comments

Comments
 (0)