Repository: incubator-airflow Updated Branches: refs/heads/master 93666f996 -> aaf308a29
[AIRFLOW-1182] SparkSubmitOperator template field Project: http://git-wip-us.apache.org/repos/asf/incubator-airflow/repo Commit: http://git-wip-us.apache.org/repos/asf/incubator-airflow/commit/f29dc7c4 Tree: http://git-wip-us.apache.org/repos/asf/incubator-airflow/tree/f29dc7c4 Diff: http://git-wip-us.apache.org/repos/asf/incubator-airflow/diff/f29dc7c4 Branch: refs/heads/master Commit: f29dc7c4c6dbc166c6bd6443e9b5c047e4d65b0c Parents: 3f546e2 Author: Vianney Foucault <[email protected]> Authored: Fri May 12 14:43:41 2017 +0200 Committer: Vianney Foucault <[email protected]> Committed: Fri May 12 14:54:48 2017 +0200 ---------------------------------------------------------------------- .../contrib/operators/spark_submit_operator.py | 3 + airflow/settings.py | 3 +- .../operators/test_spark_submit_operator.py | 90 +++++++++++++++----- 3 files changed, 74 insertions(+), 22 deletions(-) ---------------------------------------------------------------------- http://git-wip-us.apache.org/repos/asf/incubator-airflow/blob/f29dc7c4/airflow/contrib/operators/spark_submit_operator.py ---------------------------------------------------------------------- diff --git a/airflow/contrib/operators/spark_submit_operator.py b/airflow/contrib/operators/spark_submit_operator.py index 77aacd3..ca628e9 100644 --- a/airflow/contrib/operators/spark_submit_operator.py +++ b/airflow/contrib/operators/spark_submit_operator.py @@ -17,6 +17,7 @@ import logging from airflow.contrib.hooks.spark_submit_hook import SparkSubmitHook from airflow.models import BaseOperator from airflow.utils.decorators import apply_defaults +from airflow.settings import WEB_COLORS log = logging.getLogger(__name__) @@ -63,6 +64,8 @@ class SparkSubmitOperator(BaseOperator): :param verbose: Whether to pass the verbose flag to spark-submit process for debugging :type verbose: bool """ + template_fields = ('_name', '_application_args',) + ui_color = WEB_COLORS['LIGHTORANGE'] @apply_defaults def __init__(self, http://git-wip-us.apache.org/repos/asf/incubator-airflow/blob/f29dc7c4/airflow/settings.py ---------------------------------------------------------------------- diff --git a/airflow/settings.py b/airflow/settings.py index 08db96a..47e0e54 100644 --- a/airflow/settings.py +++ b/airflow/settings.py @@ -177,4 +177,5 @@ configure_orm() KILOBYTE = 1024 MEGABYTE = KILOBYTE * KILOBYTE -WEB_COLORS = {'LIGHTBLUE': '#4d9de0'} +WEB_COLORS = {'LIGHTBLUE': '#4d9de0', + 'LIGHTORANGE': '#FF9933'} http://git-wip-us.apache.org/repos/asf/incubator-airflow/blob/f29dc7c4/tests/contrib/operators/test_spark_submit_operator.py ---------------------------------------------------------------------- diff --git a/tests/contrib/operators/test_spark_submit_operator.py b/tests/contrib/operators/test_spark_submit_operator.py index 6bed6a1..09c5a93 100644 --- a/tests/contrib/operators/test_spark_submit_operator.py +++ b/tests/contrib/operators/test_spark_submit_operator.py @@ -18,6 +18,8 @@ import datetime import sys from airflow import DAG, configuration +from airflow.models import TaskInstance + from airflow.contrib.operators.spark_submit_operator import SparkSubmitOperator DEFAULT_DATE = datetime.datetime(2017, 1, 1) @@ -37,7 +39,7 @@ class TestSparkSubmitOperator(unittest.TestCase): 'executor_memory': '22g', 'keytab': 'privileged_user.keytab', 'principal': 'user/[email protected]', - 'name': 'spark-job', + 'name': '{{ task_instance.task_id }}', 'num_executors': 10, 'verbose': True, 'application': 'test_application.py', @@ -45,7 +47,9 @@ class TestSparkSubmitOperator(unittest.TestCase): 'java_class': 'com.foo.bar.AppMain', 'application_args': [ '-f foo', - '--bar bar' + '--bar bar', + '--start {{ macros.ds_add(ds, -1)}}', + '--end {{ ds }}' ] } @@ -62,33 +66,77 @@ class TestSparkSubmitOperator(unittest.TestCase): } self.dag = DAG('test_dag_id', default_args=args) - def test_execute(self, conn_id='spark_default'): + def test_execute(self): + # Given / When + conn_id = 'spark_default' operator = SparkSubmitOperator( task_id='spark_submit_job', dag=self.dag, **self._config ) - self.assertEqual(conn_id, operator._conn_id) - - self.assertEqual(self._config['application'], operator._application) - self.assertEqual(self._config['conf'], operator._conf) - self.assertEqual(self._config['files'], operator._files) - self.assertEqual(self._config['py_files'], operator._py_files) - self.assertEqual(self._config['jars'], operator._jars) - self.assertEqual(self._config['total_executor_cores'], operator._total_executor_cores) - self.assertEqual(self._config['executor_cores'], operator._executor_cores) - self.assertEqual(self._config['executor_memory'], operator._executor_memory) - self.assertEqual(self._config['keytab'], operator._keytab) - self.assertEqual(self._config['principal'], operator._principal) - self.assertEqual(self._config['name'], operator._name) - self.assertEqual(self._config['num_executors'], operator._num_executors) - self.assertEqual(self._config['verbose'], operator._verbose) - self.assertEqual(self._config['java_class'], operator._java_class) - self.assertEqual(self._config['driver_memory'], operator._driver_memory) - self.assertEqual(self._config['application_args'], operator._application_args) + # Then + expected_dict = { + 'conf': { + 'parquet.compression': 'SNAPPY' + }, + 'files': 'hive-site.xml', + 'py_files': 'sample_library.py', + 'jars': 'parquet.jar', + 'total_executor_cores': 4, + 'executor_cores': 4, + 'executor_memory': '22g', + 'keytab': 'privileged_user.keytab', + 'principal': 'user/[email protected]', + 'name': '{{ task_instance.task_id }}', + 'num_executors': 10, + 'verbose': True, + 'application': 'test_application.py', + 'driver_memory': '3g', + 'java_class': 'com.foo.bar.AppMain', + 'application_args': [ + '-f foo', + '--bar bar', + '--start {{ macros.ds_add(ds, -1)}}', + '--end {{ ds }}' + ] + } + self.assertEqual(conn_id, operator._conn_id) + self.assertEqual(expected_dict['application'], operator._application) + self.assertEqual(expected_dict['conf'], operator._conf) + self.assertEqual(expected_dict['files'], operator._files) + self.assertEqual(expected_dict['py_files'], operator._py_files) + self.assertEqual(expected_dict['jars'], operator._jars) + self.assertEqual(expected_dict['total_executor_cores'], operator._total_executor_cores) + self.assertEqual(expected_dict['executor_cores'], operator._executor_cores) + self.assertEqual(expected_dict['executor_memory'], operator._executor_memory) + self.assertEqual(expected_dict['keytab'], operator._keytab) + self.assertEqual(expected_dict['principal'], operator._principal) + self.assertEqual(expected_dict['name'], operator._name) + self.assertEqual(expected_dict['num_executors'], operator._num_executors) + self.assertEqual(expected_dict['verbose'], operator._verbose) + self.assertEqual(expected_dict['java_class'], operator._java_class) + self.assertEqual(expected_dict['driver_memory'], operator._driver_memory) + self.assertEqual(expected_dict['application_args'], operator._application_args) + + def test_render_template(self): + # Given + operator = SparkSubmitOperator(task_id='spark_submit_job', dag=self.dag, **self._config) + ti = TaskInstance(operator, DEFAULT_DATE) + + # When + ti.render_templates() + + # Then + expected_application_args = [u'-f foo', + u'--bar bar', + u'--start %s' % (DEFAULT_DATE - datetime.timedelta(days=1)).strftime("%Y-%m-%d"), + u'--end %s' % DEFAULT_DATE.strftime("%Y-%m-%d")] + expected_name = "spark_submit_job" + self.assertListEqual(sorted(expected_application_args), sorted(getattr(operator, '_application_args'))) + self.assertEqual(expected_name, getattr(operator, '_name')) if __name__ == '__main__':
