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__':

Reply via email to