From 66600c94298eb80c8f5d8fffcc62fb5d36084dff Mon Sep 17 00:00:00 2001 From: Amogh Date: Mon, 1 Apr 2024 10:13:39 +0530 Subject: [PATCH 1/6] final changes --- airflow/providers/apache/hive/hooks/hive.py | 24 +++++++++++++++++-- .../providers/apache/hive/operators/hive.py | 3 +++ 2 files changed, 25 insertions(+), 2 deletions(-) diff --git a/airflow/providers/apache/hive/hooks/hive.py b/airflow/providers/apache/hive/hooks/hive.py index cb66d362740ec..6492cce4ce0a4 100644 --- a/airflow/providers/apache/hive/hooks/hive.py +++ b/airflow/providers/apache/hive/hooks/hive.py @@ -97,9 +97,13 @@ def __init__( hive_cli_params: str = "", auth: str | None = None, proxy_user: str | None = None, + high_availability: bool | None = None, ) -> None: super().__init__() conn = self.get_connection(hive_cli_conn_id) + + print("conn is from __init__", conn) + self.hive_cli_params: str = hive_cli_params self.use_beeline: bool = conn.extra_dejson.get("use_beeline", False) self.auth = auth @@ -116,6 +120,7 @@ def __init__( self.mapred_queue_priority = mapred_queue_priority self.mapred_job_name = mapred_job_name self.proxy_user = proxy_user + self.high_availability = high_availability @classmethod def get_connection_form_widgets(cls) -> dict[str, Any]: @@ -130,6 +135,7 @@ def get_connection_form_widgets(cls) -> dict[str, Any]: "principal": StringField( lazy_gettext("Principal"), widget=BS3TextFieldWidget(), default="hive/_HOST@EXAMPLE.COM" ), + "high_availability": BooleanField(lazy_gettext("High Availability"), default=False), } @classmethod @@ -160,6 +166,10 @@ def _prepare_cli_cmd(self) -> list[Any]: hive_bin = "beeline" self._validate_beeline_parameters(conn) jdbc_url = f"jdbc:hive2://{conn.host}:{conn.port}/{conn.schema}" + print("conn is", conn) + print("conn parameters", conn.host, conn.port, conn.schema) + if self.high_availability: + jdbc_url = f"jdbc:hive2://{conn.host}/{conn.schema}" if conf.get("core", "security") == "kerberos": template = conn.extra_dejson.get("principal", "hive/_HOST@EXAMPLE.COM") if "_HOST" in template: @@ -169,7 +179,11 @@ def _prepare_cli_cmd(self) -> list[Any]: raise RuntimeError("The principal should not contain the ';' character") if ";" in proxy_user: raise RuntimeError("The proxy_user should not contain the ';' character") - jdbc_url += f";principal={template};{proxy_user}" + if proxy_user: + jdbc_url += f";principal={template};{proxy_user}" + else: + jdbc_url += f";principal={template}" \ + f";serviceDiscoveryMode=zooKeeper;ssl=true;zooKeeperNamespace=hiveserver2" elif self.auth: jdbc_url += ";auth=" + self.auth @@ -186,7 +200,13 @@ def _prepare_cli_cmd(self) -> list[Any]: return [hive_bin, *cmd_extra, *hive_params_list] def _validate_beeline_parameters(self, conn): - if ":" in conn.host or "/" in conn.host or ";" in conn.host: + if self.high_availability: + if ";" in conn.schema: + raise Exception( + f"The schema used in beeline command ({conn.schema}) should not contain ';' character)" + ) + return + elif ":" in conn.host or "/" in conn.host or ";" in conn.host: raise Exception( f"The host used in beeline command ({conn.host}) should not contain ':/;' characters)" ) diff --git a/airflow/providers/apache/hive/operators/hive.py b/airflow/providers/apache/hive/operators/hive.py index 398cadce0d829..bbe15a5342ed3 100644 --- a/airflow/providers/apache/hive/operators/hive.py +++ b/airflow/providers/apache/hive/operators/hive.py @@ -96,6 +96,7 @@ def __init__( hive_cli_params: str = "", auth: str | None = None, proxy_user: str | None = None, + high_availability: bool | None = None, **kwargs: Any, ) -> None: super().__init__(**kwargs) @@ -111,6 +112,7 @@ def __init__( self.hive_cli_params = hive_cli_params self.auth = auth self.proxy_user = proxy_user + self.high_availability = high_availability job_name_template = conf.get_mandatory_value( "hive", "mapred_job_name_template", @@ -129,6 +131,7 @@ def hook(self) -> HiveCliHook: hive_cli_params=self.hive_cli_params, auth=self.auth, proxy_user=self.proxy_user, + high_availability=self.high_availability ) @deprecated(reason="use `hook` property instead.", category=AirflowProviderDeprecationWarning) From 0f7571cefdadcd456eb9e8c2c3218e407953c25b Mon Sep 17 00:00:00 2001 From: Amogh Date: Mon, 1 Apr 2024 12:23:31 +0530 Subject: [PATCH 2/6] enhancing code --- airflow/providers/apache/hive/hooks/hive.py | 13 ++++++------- airflow/providers/apache/hive/operators/hive.py | 3 --- 2 files changed, 6 insertions(+), 10 deletions(-) diff --git a/airflow/providers/apache/hive/hooks/hive.py b/airflow/providers/apache/hive/hooks/hive.py index 6492cce4ce0a4..aeb836b788d77 100644 --- a/airflow/providers/apache/hive/hooks/hive.py +++ b/airflow/providers/apache/hive/hooks/hive.py @@ -97,7 +97,6 @@ def __init__( hive_cli_params: str = "", auth: str | None = None, proxy_user: str | None = None, - high_availability: bool | None = None, ) -> None: super().__init__() conn = self.get_connection(hive_cli_conn_id) @@ -120,7 +119,7 @@ def __init__( self.mapred_queue_priority = mapred_queue_priority self.mapred_job_name = mapred_job_name self.proxy_user = proxy_user - self.high_availability = high_availability + self.high_availability = self.conn.extra_dejson.get("high_availability", False) @classmethod def get_connection_form_widgets(cls) -> dict[str, Any]: @@ -179,11 +178,11 @@ def _prepare_cli_cmd(self) -> list[Any]: raise RuntimeError("The principal should not contain the ';' character") if ";" in proxy_user: raise RuntimeError("The proxy_user should not contain the ';' character") - if proxy_user: - jdbc_url += f";principal={template};{proxy_user}" - else: - jdbc_url += f";principal={template}" \ - f";serviceDiscoveryMode=zooKeeper;ssl=true;zooKeeperNamespace=hiveserver2" + jdbc_url += f";principal={template};{proxy_user}" + if self.high_availability: + if proxy_user: + jdbc_url += ";" + jdbc_url += "serviceDiscoveryMode=zooKeeper;ssl=true;zooKeeperNamespace=hiveserver2" elif self.auth: jdbc_url += ";auth=" + self.auth diff --git a/airflow/providers/apache/hive/operators/hive.py b/airflow/providers/apache/hive/operators/hive.py index bbe15a5342ed3..398cadce0d829 100644 --- a/airflow/providers/apache/hive/operators/hive.py +++ b/airflow/providers/apache/hive/operators/hive.py @@ -96,7 +96,6 @@ def __init__( hive_cli_params: str = "", auth: str | None = None, proxy_user: str | None = None, - high_availability: bool | None = None, **kwargs: Any, ) -> None: super().__init__(**kwargs) @@ -112,7 +111,6 @@ def __init__( self.hive_cli_params = hive_cli_params self.auth = auth self.proxy_user = proxy_user - self.high_availability = high_availability job_name_template = conf.get_mandatory_value( "hive", "mapred_job_name_template", @@ -131,7 +129,6 @@ def hook(self) -> HiveCliHook: hive_cli_params=self.hive_cli_params, auth=self.auth, proxy_user=self.proxy_user, - high_availability=self.high_availability ) @deprecated(reason="use `hook` property instead.", category=AirflowProviderDeprecationWarning) From 7e5fb6509df69d0f67fa47a275081514e3a8ef5a Mon Sep 17 00:00:00 2001 From: Amogh Date: Mon, 1 Apr 2024 14:10:03 +0530 Subject: [PATCH 3/6] adding tests --- airflow/providers/apache/hive/hooks/hive.py | 3 +- .../connections/hive_cli.rst | 4 +++ .../providers/apache/hive/hooks/test_hive.py | 34 +++++++++++++++++++ 3 files changed, 39 insertions(+), 2 deletions(-) diff --git a/airflow/providers/apache/hive/hooks/hive.py b/airflow/providers/apache/hive/hooks/hive.py index aeb836b788d77..1a0af3bc8bea3 100644 --- a/airflow/providers/apache/hive/hooks/hive.py +++ b/airflow/providers/apache/hive/hooks/hive.py @@ -165,10 +165,9 @@ def _prepare_cli_cmd(self) -> list[Any]: hive_bin = "beeline" self._validate_beeline_parameters(conn) jdbc_url = f"jdbc:hive2://{conn.host}:{conn.port}/{conn.schema}" - print("conn is", conn) - print("conn parameters", conn.host, conn.port, conn.schema) if self.high_availability: jdbc_url = f"jdbc:hive2://{conn.host}/{conn.schema}" + self.log.info("High Availability set, setting JDBC url as %s", jdbc_url) if conf.get("core", "security") == "kerberos": template = conn.extra_dejson.get("principal", "hive/_HOST@EXAMPLE.COM") if "_HOST" in template: diff --git a/docs/apache-airflow-providers-apache-hive/connections/hive_cli.rst b/docs/apache-airflow-providers-apache-hive/connections/hive_cli.rst index cc52f1db92be2..5e88df971d989 100644 --- a/docs/apache-airflow-providers-apache-hive/connections/hive_cli.rst +++ b/docs/apache-airflow-providers-apache-hive/connections/hive_cli.rst @@ -73,6 +73,10 @@ Proxy User (optional) Principal (optional) Specify the JDBC Hive principal to be used with Hive Beeline. +High Availability (optional) + Specify as ``True`` if you want to connect to a Hive installation running in high + availability mode. Specify host accordingly. + When specifying the connection in environment variable you should specify it using URI syntax. diff --git a/tests/providers/apache/hive/hooks/test_hive.py b/tests/providers/apache/hive/hooks/test_hive.py index b69f68b149632..f3bce01c895b8 100644 --- a/tests/providers/apache/hive/hooks/test_hive.py +++ b/tests/providers/apache/hive/hooks/test_hive.py @@ -901,3 +901,37 @@ def test_get_wrong_principal(self): # Run with pytest.raises(RuntimeError, match="The principal should not contain the ';' character"): hook._prepare_cli_cmd() + + @pytest.mark.parametrize( + "extra_dejson, expected_keys", + [ + ( + {"high_availability": "true"}, + "serviceDiscoveryMode=zooKeeper;ssl=true;zooKeeperNamespace=hiveserver2", + ), + ( + {"high_availability": "false"}, + "serviceDiscoveryMode=zooKeeper;ssl=true;zooKeeperNamespace=hiveserver2", + ), + ({}, "serviceDiscoveryMode=zooKeeper;ssl=true;zooKeeperNamespace=hiveserver2"), + ], + ) + def test_high_availability(self, extra_dejson, expected_keys): + hook = MockHiveCliHook() + returner = mock.MagicMock() + returner.extra_dejson = extra_dejson + returner.login = "admin" + hook.use_beeline = True + hook.conn = returner + hook.high_availability = ( + True + if ("high_availability" in extra_dejson and extra_dejson["high_availability"] == "true") + else False + ) + + result = hook._prepare_cli_cmd() + + if hook.high_availability: + assert expected_keys in result[2] + else: + assert expected_keys not in result[2] From 7d2f771a10ed9ce1bfd2427f623e4b6b86d30048 Mon Sep 17 00:00:00 2001 From: Amogh Date: Mon, 1 Apr 2024 16:31:48 +0530 Subject: [PATCH 4/6] review comments from romsharon98 --- tests/providers/apache/hive/hooks/test_hive.py | 6 ++++++ 1 file changed, 6 insertions(+) diff --git a/tests/providers/apache/hive/hooks/test_hive.py b/tests/providers/apache/hive/hooks/test_hive.py index f3bce01c895b8..e444697bf69bd 100644 --- a/tests/providers/apache/hive/hooks/test_hive.py +++ b/tests/providers/apache/hive/hooks/test_hive.py @@ -914,6 +914,12 @@ def test_get_wrong_principal(self): "serviceDiscoveryMode=zooKeeper;ssl=true;zooKeeperNamespace=hiveserver2", ), ({}, "serviceDiscoveryMode=zooKeeper;ssl=true;zooKeeperNamespace=hiveserver2"), + # with proxy user + ( + {"proxy_user": "a_user_proxy", "high_availability": "true"}, + "hive.server2.proxy.user=a_user_proxy;" + "serviceDiscoveryMode=zooKeeper;ssl=true;zooKeeperNamespace=hiveserver2", + ), ], ) def test_high_availability(self, extra_dejson, expected_keys): From e855f9da8727800ad2f5f183d76484f1ce108d3c Mon Sep 17 00:00:00 2001 From: Amogh Date: Tue, 2 Apr 2024 09:48:38 +0530 Subject: [PATCH 5/6] removing print statement --- airflow/providers/apache/hive/hooks/hive.py | 3 --- 1 file changed, 3 deletions(-) diff --git a/airflow/providers/apache/hive/hooks/hive.py b/airflow/providers/apache/hive/hooks/hive.py index 1a0af3bc8bea3..793f81bd56399 100644 --- a/airflow/providers/apache/hive/hooks/hive.py +++ b/airflow/providers/apache/hive/hooks/hive.py @@ -100,9 +100,6 @@ def __init__( ) -> None: super().__init__() conn = self.get_connection(hive_cli_conn_id) - - print("conn is from __init__", conn) - self.hive_cli_params: str = hive_cli_params self.use_beeline: bool = conn.extra_dejson.get("use_beeline", False) self.auth = auth From 2dfcd1d6c382f84f82c0b858b770dc17e6dbd804 Mon Sep 17 00:00:00 2001 From: Amogh Date: Wed, 3 Apr 2024 14:43:00 +0530 Subject: [PATCH 6/6] review comments from husseinawala --- airflow/providers/apache/hive/hooks/hive.py | 6 ++++-- 1 file changed, 4 insertions(+), 2 deletions(-) diff --git a/airflow/providers/apache/hive/hooks/hive.py b/airflow/providers/apache/hive/hooks/hive.py index 793f81bd56399..cd508b21866ce 100644 --- a/airflow/providers/apache/hive/hooks/hive.py +++ b/airflow/providers/apache/hive/hooks/hive.py @@ -161,10 +161,12 @@ def _prepare_cli_cmd(self) -> list[Any]: if self.use_beeline: hive_bin = "beeline" self._validate_beeline_parameters(conn) - jdbc_url = f"jdbc:hive2://{conn.host}:{conn.port}/{conn.schema}" if self.high_availability: jdbc_url = f"jdbc:hive2://{conn.host}/{conn.schema}" self.log.info("High Availability set, setting JDBC url as %s", jdbc_url) + else: + jdbc_url = f"jdbc:hive2://{conn.host}:{conn.port}/{conn.schema}" + self.log.info("High Availability not set, setting JDBC url as %s", jdbc_url) if conf.get("core", "security") == "kerberos": template = conn.extra_dejson.get("principal", "hive/_HOST@EXAMPLE.COM") if "_HOST" in template: @@ -176,7 +178,7 @@ def _prepare_cli_cmd(self) -> list[Any]: raise RuntimeError("The proxy_user should not contain the ';' character") jdbc_url += f";principal={template};{proxy_user}" if self.high_availability: - if proxy_user: + if not jdbc_url.endswith(";"): jdbc_url += ";" jdbc_url += "serviceDiscoveryMode=zooKeeper;ssl=true;zooKeeperNamespace=hiveserver2" elif self.auth: