From c1a3b5f97ca1095bdc12c8d3420c0bf088b3b9ab Mon Sep 17 00:00:00 2001 From: Karun Poudel Date: Mon, 15 Sep 2025 17:39:41 +0000 Subject: [PATCH 01/17] init fx fx pre-commit progress fx fx fx --- .../postgres/docs/connections/postgres.rst | 7 +++- .../providers/postgres/hooks/postgres.py | 34 +++++++++++++++---- .../unit/postgres/hooks/test_postgres.py | 27 +++++++++++++++ 3 files changed, 61 insertions(+), 7 deletions(-) diff --git a/providers/postgres/docs/connections/postgres.rst b/providers/postgres/docs/connections/postgres.rst index 539620ad08cc9..e96706ad458e8 100644 --- a/providers/postgres/docs/connections/postgres.rst +++ b/providers/postgres/docs/connections/postgres.rst @@ -96,7 +96,9 @@ Extra (optional) * ``iam`` - If set to ``True`` than use AWS IAM database authentication for `Amazon RDS `__, `Amazon Aurora `__ - or `Amazon Redshift `__. + `Amazon Redshift `__ + or use Microsoft Entra Authentication for + `Azure Postgres Flexible Server `__. * ``aws_conn_id`` - AWS Connection ID which use for authentication via AWS IAM, if not specified then **aws_default** is used. * ``redshift`` - Used when AWS IAM database authentication enabled. @@ -104,6 +106,9 @@ Extra (optional) * ``cluster-identifier`` - The unique identifier of the Amazon Redshift Cluster that contains the database for which you are requesting credentials. This parameter is case sensitive. If not specified than hostname from **Connection Host** is used. + * ``azure_conn_id`` - Azure Connection ID to be used for authentication via Azure Entra ID. Azure Oauth token + is retrieved from the azure connection which is used as password for PostgreSQL connection. + Example "extras" field (Amazon RDS PostgreSQL or Amazon Aurora PostgreSQL): diff --git a/providers/postgres/src/airflow/providers/postgres/hooks/postgres.py b/providers/postgres/src/airflow/providers/postgres/hooks/postgres.py index 7a3be0ff4e314..65082db29ebdc 100644 --- a/providers/postgres/src/airflow/providers/postgres/hooks/postgres.py +++ b/providers/postgres/src/airflow/providers/postgres/hooks/postgres.py @@ -37,6 +37,11 @@ from airflow.providers.common.sql.hooks.sql import DbApiHook from airflow.providers.postgres.dialects.postgres import PostgresDialect +try: + from airflow.sdk import Connection +except ImportError: + from airflow.models.connection import Connection # type: ignore[assignment] + USE_PSYCOPG3: bool try: import psycopg as psycopg # needed for patching in unit tests @@ -64,11 +69,6 @@ if USE_PSYCOPG3: from psycopg.errors import Diagnostic - try: - from airflow.sdk import Connection - except ImportError: - from airflow.models.connection import Connection # type: ignore[assignment] - CursorType: TypeAlias = DictCursor | RealDictCursor | NamedTupleCursor CursorRow: TypeAlias = dict[str, Any] | tuple[Any, ...] @@ -156,7 +156,9 @@ class PostgresHook(DbApiHook): "aws_conn_id", "sqlalchemy_scheme", "sqlalchemy_query", + "azure_conn_id", } + azure_oauth_scope = "https://ossrdbms-aad.database.windows.net/.default" def __init__( self, *args, options: str | None = None, enable_log_db_messages: bool = False, **kwargs @@ -177,6 +179,8 @@ def sqlalchemy_url(self) -> URL: query = conn.extra_dejson.get("sqlalchemy_query", {}) if not isinstance(query, dict): raise AirflowException("The parameter 'sqlalchemy_query' must be of type dict!") + if conn.extra_dejson.get("iam", False): + conn.login, conn.password, conn.port = self.get_iam_token(conn) return URL.create( drivername="postgresql+psycopg" if USE_PSYCOPG3 else "postgresql", username=self.__cast_nullable(conn.login, str), @@ -441,8 +445,14 @@ def _serialize_cell(cell: object, conn: Any | None = None) -> Any: return PostgresHook._serialize_cell_ppg2(cell, conn) def get_iam_token(self, conn: Connection) -> tuple[str, str, int]: + """Get the IAM token from different identity providers.""" + if conn.extra_dejson.get("azure_conn_id"): + return self.get_azure_iam_token(conn) + return self.get_aws_iam_token(conn) + + def get_aws_iam_token(self, conn: Connection) -> tuple[str, str, int]: """ - Get the IAM token. + Get the AWS IAM token. This uses AWSHook to retrieve a temporary password to connect to Postgres or Redshift. Port is required. If none is provided, the default @@ -500,6 +510,18 @@ def get_iam_token(self, conn: Connection) -> tuple[str, str, int]: token = rds_client.generate_db_auth_token(conn.host, port, conn.login) return cast("str", login), cast("str", token), port + def get_azure_iam_token(self, conn: Connection) -> tuple[str, str, int]: + """ + Get the Azure IAM token. + + This uses AzureBaseHook to retrieve an OAUTH token to connect to Postgres. + """ + azure_conn_id = conn.extra_dejson.get("azure_conn_id", "azure_default") + azure_conn = Connection.get(azure_conn_id) + azure_base_hook = azure_conn.get_hook() + token = azure_base_hook.get_token(self.azure_oauth_scope).token + return cast("str", conn.login or azure_conn.login), cast("str", token), conn.port or 5432 + def get_table_primary_key(self, table: str, schema: str | None = "public") -> list[str] | None: """ Get the table's primary key. diff --git a/providers/postgres/tests/unit/postgres/hooks/test_postgres.py b/providers/postgres/tests/unit/postgres/hooks/test_postgres.py index f394c561c2c53..abd983c495f3f 100644 --- a/providers/postgres/tests/unit/postgres/hooks/test_postgres.py +++ b/providers/postgres/tests/unit/postgres/hooks/test_postgres.py @@ -444,6 +444,33 @@ def test_get_conn_rds_iam_redshift_serverless( port=(port or 5439), ) + def test_get_conn_azure_iam(self, mocker, mock_connect): + mock_azure_conn_id = "azure_conn1" + mock_db_token = "azure_token1" + mock_conn_extra = {"iam": True, "azure_conn_id": mock_azure_conn_id} + self.connection.extra = json.dumps(mock_conn_extra) + + mock_connection_class = mocker.patch("airflow.providers.postgres.hooks.postgres.Connection") + mock_azure_base_hook = mock_connection_class.get.return_value.get_hook.return_value + mock_azure_base_hook.get_token.return_value.token = mock_db_token + + self.db_hook.get_conn() + + # Check AzureBaseHook initialization and get_token call args + mock_connection_class.get.assert_called_once_with(mock_azure_conn_id) + mock_azure_base_hook.get_token.assert_called_once_with(PostgresHook.azure_oauth_scope) + + # Check expected psycopg2 connection call args + mock_connect.assert_called_once_with( + user=self.connection.login, + password=mock_db_token, + host=self.connection.host, + dbname=self.connection.schema, + port=(self.connection.port or 5432), + ) + + assert mock_db_token in self.db_hook.sqlalchemy_url + def test_get_uri_from_connection_without_database_override(self, mocker): expected: str = f"postgresql{'+psycopg' if USE_PSYCOPG3 else ''}://login:password@host:1/database" self.db_hook.get_connection = mocker.MagicMock( From 7a9914f5705d9166af7011abcb95a2dbf0667521 Mon Sep 17 00:00:00 2001 From: Karun Poudel Date: Tue, 16 Sep 2025 09:23:37 +0000 Subject: [PATCH 02/17] fx --- providers/postgres/docs/connections/postgres.rst | 9 +++++++++ providers/postgres/provider.yaml | 13 +++++++++++++ 2 files changed, 22 insertions(+) diff --git a/providers/postgres/docs/connections/postgres.rst b/providers/postgres/docs/connections/postgres.rst index e96706ad458e8..2a009cce8729d 100644 --- a/providers/postgres/docs/connections/postgres.rst +++ b/providers/postgres/docs/connections/postgres.rst @@ -130,6 +130,15 @@ Extra (optional) "cluster-identifier": "awesome-redshift-identifier" } + Example "extras" field (to use Azure Entra Authentication for Postgres Flexible Server): + + .. code-block:: json + + { + "iam": true, + "azure_conn_id": "azure_default_conn" + } + When specifying the connection as URI (in :envvar:`AIRFLOW_CONN_{CONN_ID}` variable) you should specify it following the standard syntax of DB connections, where extras are passed as parameters of the URI (note that all components of the URI should be URL-encoded). diff --git a/providers/postgres/provider.yaml b/providers/postgres/provider.yaml index 01fa0ee058b23..19bace1563d96 100644 --- a/providers/postgres/provider.yaml +++ b/providers/postgres/provider.yaml @@ -109,3 +109,16 @@ asset-uris: dataset-uris: - schemes: [postgres, postgresql] handler: airflow.providers.postgres.assets.postgres.sanitize_uri + +config: + postgres: + description: | + Configuration for Postgres hooks and operators. + options: + azure_oauth_scope: + description: | + The scope to use while retrieving Oauth token for Postgres Flexible Server from Azure Entra authentication. + version_added: 6.3.1 + type: string + example: ~ + default: "https://ossrdbms-aad.database.windows.net/.default" From 13343fd199cfceb6f3d9118c640d1b92d049eef6 Mon Sep 17 00:00:00 2001 From: Karun Poudel Date: Tue, 16 Sep 2025 15:27:35 +0000 Subject: [PATCH 03/17] add config --- .../src/airflow/providers/postgres/hooks/postgres.py | 6 ++++-- .../postgres/tests/unit/postgres/hooks/test_postgres.py | 2 +- 2 files changed, 5 insertions(+), 3 deletions(-) diff --git a/providers/postgres/src/airflow/providers/postgres/hooks/postgres.py b/providers/postgres/src/airflow/providers/postgres/hooks/postgres.py index 65082db29ebdc..43f0bb426bfb5 100644 --- a/providers/postgres/src/airflow/providers/postgres/hooks/postgres.py +++ b/providers/postgres/src/airflow/providers/postgres/hooks/postgres.py @@ -30,6 +30,7 @@ from psycopg2.extras import DictCursor, NamedTupleCursor, RealDictCursor, execute_batch from sqlalchemy.engine import URL +from airflow.configuration import conf from airflow.exceptions import ( AirflowException, AirflowOptionalProviderFeatureException, @@ -158,7 +159,7 @@ class PostgresHook(DbApiHook): "sqlalchemy_query", "azure_conn_id", } - azure_oauth_scope = "https://ossrdbms-aad.database.windows.net/.default" + default_azure_oauth_scope = "https://ossrdbms-aad.database.windows.net/.default" def __init__( self, *args, options: str | None = None, enable_log_db_messages: bool = False, **kwargs @@ -519,7 +520,8 @@ def get_azure_iam_token(self, conn: Connection) -> tuple[str, str, int]: azure_conn_id = conn.extra_dejson.get("azure_conn_id", "azure_default") azure_conn = Connection.get(azure_conn_id) azure_base_hook = azure_conn.get_hook() - token = azure_base_hook.get_token(self.azure_oauth_scope).token + scope = conf.get("postgres", "azure_oauth_scope", fallback=self.default_azure_oauth_scope) + token = azure_base_hook.get_token(scope).token return cast("str", conn.login or azure_conn.login), cast("str", token), conn.port or 5432 def get_table_primary_key(self, table: str, schema: str | None = "public") -> list[str] | None: diff --git a/providers/postgres/tests/unit/postgres/hooks/test_postgres.py b/providers/postgres/tests/unit/postgres/hooks/test_postgres.py index abd983c495f3f..7a8f1acd71b87 100644 --- a/providers/postgres/tests/unit/postgres/hooks/test_postgres.py +++ b/providers/postgres/tests/unit/postgres/hooks/test_postgres.py @@ -458,7 +458,7 @@ def test_get_conn_azure_iam(self, mocker, mock_connect): # Check AzureBaseHook initialization and get_token call args mock_connection_class.get.assert_called_once_with(mock_azure_conn_id) - mock_azure_base_hook.get_token.assert_called_once_with(PostgresHook.azure_oauth_scope) + mock_azure_base_hook.get_token.assert_called_once_with(PostgresHook.default_azure_oauth_scope) # Check expected psycopg2 connection call args mock_connect.assert_called_once_with( From 7dedc8a99815011a6bf164c5fe433a5815806dc1 Mon Sep 17 00:00:00 2001 From: Karun Poudel Date: Tue, 16 Sep 2025 16:08:53 +0000 Subject: [PATCH 04/17] conf doc --- .../postgres/docs/configurations-ref.rst | 19 +++++++++++++++++++ providers/postgres/docs/index.rst | 1 + 2 files changed, 20 insertions(+) create mode 100644 providers/postgres/docs/configurations-ref.rst diff --git a/providers/postgres/docs/configurations-ref.rst b/providers/postgres/docs/configurations-ref.rst new file mode 100644 index 0000000000000..ea8e668d75793 --- /dev/null +++ b/providers/postgres/docs/configurations-ref.rst @@ -0,0 +1,19 @@ + .. Licensed to the Apache Software Foundation (ASF) under one + or more contributor license agreements. See the NOTICE file + distributed with this work for additional information + regarding copyright ownership. The ASF licenses this file + to you under the Apache License, Version 2.0 (the + "License"); you may not use this file except in compliance + with the License. You may obtain a copy of the License at + + .. http://www.apache.org/licenses/LICENSE-2.0 + + .. Unless required by applicable law or agreed to in writing, + software distributed under the License is distributed on an + "AS IS" BASIS, WITHOUT WARRANTIES OR CONDITIONS OF ANY + KIND, either express or implied. See the License for the + specific language governing permissions and limitations + under the License. + +.. include:: /../../../../devel-common/src/sphinx_exts/includes/providers-configurations-ref.rst +.. include:: /../../../../devel-common/src/sphinx_exts/includes/sections-and-options.rst diff --git a/providers/postgres/docs/index.rst b/providers/postgres/docs/index.rst index 800bf10d57e43..54535c8d936da 100644 --- a/providers/postgres/docs/index.rst +++ b/providers/postgres/docs/index.rst @@ -41,6 +41,7 @@ :maxdepth: 1 :caption: References + Configuration Python API <_api/airflow/providers/postgres/index> Dialects From fc60d7d9cba3d2e967071b2e45671e610b0a1e4e Mon Sep 17 00:00:00 2001 From: Karun Poudel Date: Tue, 16 Sep 2025 16:11:31 +0000 Subject: [PATCH 05/17] fx --- providers/postgres/docs/connections/postgres.rst | 1 - 1 file changed, 1 deletion(-) diff --git a/providers/postgres/docs/connections/postgres.rst b/providers/postgres/docs/connections/postgres.rst index 2a009cce8729d..46d07991ec15a 100644 --- a/providers/postgres/docs/connections/postgres.rst +++ b/providers/postgres/docs/connections/postgres.rst @@ -109,7 +109,6 @@ Extra (optional) * ``azure_conn_id`` - Azure Connection ID to be used for authentication via Azure Entra ID. Azure Oauth token is retrieved from the azure connection which is used as password for PostgreSQL connection. - Example "extras" field (Amazon RDS PostgreSQL or Amazon Aurora PostgreSQL): .. code-block:: json From 93e48488fa1a4d8fc49c5cea8feac10c8c747c98 Mon Sep 17 00:00:00 2001 From: Karun Poudel Date: Tue, 16 Sep 2025 18:36:36 +0000 Subject: [PATCH 06/17] fix1 --- docs/spelling_wordlist.txt | 1 + .../providers/postgres/get_provider_info.py | 14 ++++++++++++++ 2 files changed, 15 insertions(+) diff --git a/docs/spelling_wordlist.txt b/docs/spelling_wordlist.txt index 5163b2f91039e..9d21cb7a2bf19 100644 --- a/docs/spelling_wordlist.txt +++ b/docs/spelling_wordlist.txt @@ -607,6 +607,7 @@ encodable encryptor enqueue enqueued +Entra Entry EntryGroup EntryGroups diff --git a/providers/postgres/src/airflow/providers/postgres/get_provider_info.py b/providers/postgres/src/airflow/providers/postgres/get_provider_info.py index e33bc651039fe..55513bdc99513 100644 --- a/providers/postgres/src/airflow/providers/postgres/get_provider_info.py +++ b/providers/postgres/src/airflow/providers/postgres/get_provider_info.py @@ -65,4 +65,18 @@ def get_provider_info(): "handler": "airflow.providers.postgres.assets.postgres.sanitize_uri", } ], + "config": { + "postgres": { + "description": "Configuration for Postgres hooks and operators.\n", + "options": { + "azure_oauth_scope": { + "description": "The scope to use while retrieving Oauth token for Postgres Flexible Server from Azure Entra authentication.\n", + "version_added": "6.3.1", + "type": "string", + "example": None, + "default": "https://ossrdbms-aad.database.windows.net/.default", + } + }, + } + }, } From b4acfd8e669bc78e6ec4b0001b5154cf3257a49d Mon Sep 17 00:00:00 2001 From: Karun Poudel Date: Tue, 16 Sep 2025 19:00:02 +0000 Subject: [PATCH 07/17] update version --- providers/postgres/provider.yaml | 2 +- .../src/airflow/providers/postgres/get_provider_info.py | 2 +- 2 files changed, 2 insertions(+), 2 deletions(-) diff --git a/providers/postgres/provider.yaml b/providers/postgres/provider.yaml index 19bace1563d96..962586f2c4dc2 100644 --- a/providers/postgres/provider.yaml +++ b/providers/postgres/provider.yaml @@ -118,7 +118,7 @@ config: azure_oauth_scope: description: | The scope to use while retrieving Oauth token for Postgres Flexible Server from Azure Entra authentication. - version_added: 6.3.1 + version_added: 6.4.0 type: string example: ~ default: "https://ossrdbms-aad.database.windows.net/.default" diff --git a/providers/postgres/src/airflow/providers/postgres/get_provider_info.py b/providers/postgres/src/airflow/providers/postgres/get_provider_info.py index 55513bdc99513..c9e944d4a7b16 100644 --- a/providers/postgres/src/airflow/providers/postgres/get_provider_info.py +++ b/providers/postgres/src/airflow/providers/postgres/get_provider_info.py @@ -71,7 +71,7 @@ def get_provider_info(): "options": { "azure_oauth_scope": { "description": "The scope to use while retrieving Oauth token for Postgres Flexible Server from Azure Entra authentication.\n", - "version_added": "6.3.1", + "version_added": "6.4.0", "type": "string", "example": None, "default": "https://ossrdbms-aad.database.windows.net/.default", From 491880b679d90cd6087fa78f8445534f38e4e379 Mon Sep 17 00:00:00 2001 From: Karun Poudel Date: Fri, 19 Sep 2025 04:12:25 +0000 Subject: [PATCH 08/17] fix doc --- providers/postgres/docs/configurations-ref.rst | 4 ++-- 1 file changed, 2 insertions(+), 2 deletions(-) diff --git a/providers/postgres/docs/configurations-ref.rst b/providers/postgres/docs/configurations-ref.rst index ea8e668d75793..a52b21b2e5679 100644 --- a/providers/postgres/docs/configurations-ref.rst +++ b/providers/postgres/docs/configurations-ref.rst @@ -15,5 +15,5 @@ specific language governing permissions and limitations under the License. -.. include:: /../../../../devel-common/src/sphinx_exts/includes/providers-configurations-ref.rst -.. include:: /../../../../devel-common/src/sphinx_exts/includes/sections-and-options.rst +.. include:: /../../../devel-common/src/sphinx_exts/includes/providers-configurations-ref.rst +.. include:: /../../../devel-common/src/sphinx_exts/includes/sections-and-options.rst From b556f92dfa66fb5d1f81df052a1cced03d4d02ae Mon Sep 17 00:00:00 2001 From: Karun Poudel Date: Fri, 19 Sep 2025 06:47:54 +0000 Subject: [PATCH 09/17] prek fix --- .../src/airflow/providers/postgres/get_provider_info.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/providers/postgres/src/airflow/providers/postgres/get_provider_info.py b/providers/postgres/src/airflow/providers/postgres/get_provider_info.py index c9e944d4a7b16..3c8737cff1467 100644 --- a/providers/postgres/src/airflow/providers/postgres/get_provider_info.py +++ b/providers/postgres/src/airflow/providers/postgres/get_provider_info.py @@ -78,5 +78,5 @@ def get_provider_info(): } }, } - }, + }, } From 092fa66f4f6d8b5334e15cd18c9e0f1cf7845822 Mon Sep 17 00:00:00 2001 From: Karun Poudel Date: Fri, 3 Oct 2025 23:14:50 +0000 Subject: [PATCH 10/17] updates --- providers/postgres/docs/connections/postgres.rst | 2 +- providers/postgres/docs/index.rst | 15 ++++++++------- providers/postgres/provider.yaml | 3 ++- providers/postgres/pyproject.toml | 4 ++++ .../providers/postgres/get_provider_info.py | 2 +- .../airflow/providers/postgres/hooks/postgres.py | 15 +++++++++++++-- .../tests/unit/postgres/hooks/test_postgres.py | 2 +- 7 files changed, 30 insertions(+), 13 deletions(-) diff --git a/providers/postgres/docs/connections/postgres.rst b/providers/postgres/docs/connections/postgres.rst index 46d07991ec15a..2583c199f6ed5 100644 --- a/providers/postgres/docs/connections/postgres.rst +++ b/providers/postgres/docs/connections/postgres.rst @@ -107,7 +107,7 @@ Extra (optional) for which you are requesting credentials. This parameter is case sensitive. If not specified than hostname from **Connection Host** is used. * ``azure_conn_id`` - Azure Connection ID to be used for authentication via Azure Entra ID. Azure Oauth token - is retrieved from the azure connection which is used as password for PostgreSQL connection. + is retrieved from the azure connection which is used as password for PostgreSQL connection. Requires `apache-airflow-providers-microsoft-azure>=12.8.0`. Example "extras" field (Amazon RDS PostgreSQL or Amazon Aurora PostgreSQL): diff --git a/providers/postgres/docs/index.rst b/providers/postgres/docs/index.rst index 54535c8d936da..54953989634f7 100644 --- a/providers/postgres/docs/index.rst +++ b/providers/postgres/docs/index.rst @@ -121,13 +121,14 @@ You can install such cross-provider dependencies when installing from PyPI. For pip install apache-airflow-providers-postgres[amazon] -============================================================================================================== =============== -Dependent package Extra -============================================================================================================== =============== -`apache-airflow-providers-amazon `_ ``amazon`` -`apache-airflow-providers-common-sql `_ ``common.sql`` -`apache-airflow-providers-openlineage `_ ``openlineage`` -============================================================================================================== =============== +====================================================================================================================== =============== +Dependent package Extra +====================================================================================================================== =============== +`apache-airflow-providers-amazon `_ ``amazon`` +`apache-airflow-providers-common-sql `_ ``common.sql`` +`apache-airflow-providers-openlineage `_ ``openlineage`` +`apache-airflow-providers-microsoft-azure `_ ``microsoft.azure`` +====================================================================================================================== =============== Downloading official packages ----------------------------- diff --git a/providers/postgres/provider.yaml b/providers/postgres/provider.yaml index 962586f2c4dc2..97c0397015ad9 100644 --- a/providers/postgres/provider.yaml +++ b/providers/postgres/provider.yaml @@ -117,7 +117,8 @@ config: options: azure_oauth_scope: description: | - The scope to use while retrieving Oauth token for Postgres Flexible Server from Azure Entra authentication. + The scope to use while retrieving Oauth token for Postgres Flexible Server + from Azure Entra authentication. version_added: 6.4.0 type: string example: ~ diff --git a/providers/postgres/pyproject.toml b/providers/postgres/pyproject.toml index 5c27f689aeeef..105a5812b1f79 100644 --- a/providers/postgres/pyproject.toml +++ b/providers/postgres/pyproject.toml @@ -70,6 +70,9 @@ dependencies = [ "amazon" = [ "apache-airflow-providers-amazon>=2.6.0", ] +"microsoft.azure" = [ + "apache-airflow-providers-microsoft-azure>=12.8.0" +] "openlineage" = [ "apache-airflow-providers-openlineage" ] @@ -91,6 +94,7 @@ dev = [ "apache-airflow-devel-common", "apache-airflow-providers-amazon", "apache-airflow-providers-common-sql", + "apache-airflow-providers-microsoft-azure", "apache-airflow-providers-openlineage", # Additional devel dependencies (do not remove this line and add extra development dependencies) "apache-airflow-providers-common-sql[pandas]", diff --git a/providers/postgres/src/airflow/providers/postgres/get_provider_info.py b/providers/postgres/src/airflow/providers/postgres/get_provider_info.py index 3c8737cff1467..ba50c431a1f2d 100644 --- a/providers/postgres/src/airflow/providers/postgres/get_provider_info.py +++ b/providers/postgres/src/airflow/providers/postgres/get_provider_info.py @@ -70,7 +70,7 @@ def get_provider_info(): "description": "Configuration for Postgres hooks and operators.\n", "options": { "azure_oauth_scope": { - "description": "The scope to use while retrieving Oauth token for Postgres Flexible Server from Azure Entra authentication.\n", + "description": "The scope to use while retrieving Oauth token for Postgres Flexible Server\nfrom Azure Entra authentication.\n", "version_added": "6.4.0", "type": "string", "example": None, diff --git a/providers/postgres/src/airflow/providers/postgres/hooks/postgres.py b/providers/postgres/src/airflow/providers/postgres/hooks/postgres.py index 43f0bb426bfb5..9ff25832b360d 100644 --- a/providers/postgres/src/airflow/providers/postgres/hooks/postgres.py +++ b/providers/postgres/src/airflow/providers/postgres/hooks/postgres.py @@ -518,10 +518,21 @@ def get_azure_iam_token(self, conn: Connection) -> tuple[str, str, int]: This uses AzureBaseHook to retrieve an OAUTH token to connect to Postgres. """ azure_conn_id = conn.extra_dejson.get("azure_conn_id", "azure_default") - azure_conn = Connection.get(azure_conn_id) + try: + azure_conn = Connection.get(azure_conn_id) + except AttributeError: + azure_conn = Connection.get_connection_from_secrets(azure_conn_id) # type: ignore[attr-defined] azure_base_hook = azure_conn.get_hook() scope = conf.get("postgres", "azure_oauth_scope", fallback=self.default_azure_oauth_scope) - token = azure_base_hook.get_token(scope).token + try: + token = azure_base_hook.get_token(scope).token + except AttributeError as e: + if "get_token" in str(e): + raise AttributeError( + "'AzureBaseHook' object has no attribute 'get_token'. " + "Please upgrade apache-airflow-providers-microsoft-azure>=12.8.0." + ) from e + raise return cast("str", conn.login or azure_conn.login), cast("str", token), conn.port or 5432 def get_table_primary_key(self, table: str, schema: str | None = "public") -> list[str] | None: diff --git a/providers/postgres/tests/unit/postgres/hooks/test_postgres.py b/providers/postgres/tests/unit/postgres/hooks/test_postgres.py index eaf38a7128079..3600aa9088831 100644 --- a/providers/postgres/tests/unit/postgres/hooks/test_postgres.py +++ b/providers/postgres/tests/unit/postgres/hooks/test_postgres.py @@ -452,7 +452,7 @@ def test_get_conn_azure_iam(self, mocker, mock_connect): mock_connection_class = mocker.patch("airflow.providers.postgres.hooks.postgres.Connection") mock_azure_base_hook = mock_connection_class.get.return_value.get_hook.return_value - mock_azure_base_hook.get_token.return_value.token = mock_db_token + mock_azure_base_hook.get_token.return_value.token = "abc" self.db_hook.get_conn() From 4cc0acb29af51799f4d0c0136d9b965551baba10 Mon Sep 17 00:00:00 2001 From: Karun Poudel Date: Fri, 3 Oct 2025 23:37:15 +0000 Subject: [PATCH 11/17] fix --- providers/postgres/tests/unit/postgres/hooks/test_postgres.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/providers/postgres/tests/unit/postgres/hooks/test_postgres.py b/providers/postgres/tests/unit/postgres/hooks/test_postgres.py index 3600aa9088831..eaf38a7128079 100644 --- a/providers/postgres/tests/unit/postgres/hooks/test_postgres.py +++ b/providers/postgres/tests/unit/postgres/hooks/test_postgres.py @@ -452,7 +452,7 @@ def test_get_conn_azure_iam(self, mocker, mock_connect): mock_connection_class = mocker.patch("airflow.providers.postgres.hooks.postgres.Connection") mock_azure_base_hook = mock_connection_class.get.return_value.get_hook.return_value - mock_azure_base_hook.get_token.return_value.token = "abc" + mock_azure_base_hook.get_token.return_value.token = mock_db_token self.db_hook.get_conn() From 0956c368d8a838e238a8e0c757acd029081e1383 Mon Sep 17 00:00:00 2001 From: karunpoudel <62040859+karunpoudel@users.noreply.github.com> Date: Fri, 3 Oct 2025 20:06:00 -0400 Subject: [PATCH 12/17] Update pyproject.toml --- providers/postgres/pyproject.toml | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/providers/postgres/pyproject.toml b/providers/postgres/pyproject.toml index 105a5812b1f79..6a2e1aee0ab7c 100644 --- a/providers/postgres/pyproject.toml +++ b/providers/postgres/pyproject.toml @@ -71,7 +71,7 @@ dependencies = [ "apache-airflow-providers-amazon>=2.6.0", ] "microsoft.azure" = [ - "apache-airflow-providers-microsoft-azure>=12.8.0" + "apache-airflow-providers-microsoft-azure" ] "openlineage" = [ "apache-airflow-providers-openlineage" From d919a8dadb32d561cb4ccee1b12085bd23a901c0 Mon Sep 17 00:00:00 2001 From: karunpoudel <62040859+karunpoudel@users.noreply.github.com> Date: Sat, 4 Oct 2025 01:01:10 -0400 Subject: [PATCH 13/17] Update pyproject.toml --- providers/postgres/pyproject.toml | 1 - 1 file changed, 1 deletion(-) diff --git a/providers/postgres/pyproject.toml b/providers/postgres/pyproject.toml index 6a2e1aee0ab7c..1395be3fbd402 100644 --- a/providers/postgres/pyproject.toml +++ b/providers/postgres/pyproject.toml @@ -94,7 +94,6 @@ dev = [ "apache-airflow-devel-common", "apache-airflow-providers-amazon", "apache-airflow-providers-common-sql", - "apache-airflow-providers-microsoft-azure", "apache-airflow-providers-openlineage", # Additional devel dependencies (do not remove this line and add extra development dependencies) "apache-airflow-providers-common-sql[pandas]", From ee8e349c3125b3a6d4e570faa9587fea0b3ba21b Mon Sep 17 00:00:00 2001 From: Karun Poudel Date: Sat, 4 Oct 2025 06:02:13 +0000 Subject: [PATCH 14/17] fix --- providers/postgres/pyproject.toml | 1 + .../src/airflow/providers/postgres/hooks/postgres.py | 7 +++++-- 2 files changed, 6 insertions(+), 2 deletions(-) diff --git a/providers/postgres/pyproject.toml b/providers/postgres/pyproject.toml index 1395be3fbd402..6a2e1aee0ab7c 100644 --- a/providers/postgres/pyproject.toml +++ b/providers/postgres/pyproject.toml @@ -94,6 +94,7 @@ dev = [ "apache-airflow-devel-common", "apache-airflow-providers-amazon", "apache-airflow-providers-common-sql", + "apache-airflow-providers-microsoft-azure", "apache-airflow-providers-openlineage", # Additional devel dependencies (do not remove this line and add extra development dependencies) "apache-airflow-providers-common-sql[pandas]", diff --git a/providers/postgres/src/airflow/providers/postgres/hooks/postgres.py b/providers/postgres/src/airflow/providers/postgres/hooks/postgres.py index 9ff25832b360d..622c5960a926c 100644 --- a/providers/postgres/src/airflow/providers/postgres/hooks/postgres.py +++ b/providers/postgres/src/airflow/providers/postgres/hooks/postgres.py @@ -517,12 +517,15 @@ def get_azure_iam_token(self, conn: Connection) -> tuple[str, str, int]: This uses AzureBaseHook to retrieve an OAUTH token to connect to Postgres. """ + if TYPE_CHECKING: + from airflow.providers.microsoft.azure.hooks.base_azure import AzureBaseHook + azure_conn_id = conn.extra_dejson.get("azure_conn_id", "azure_default") try: azure_conn = Connection.get(azure_conn_id) except AttributeError: azure_conn = Connection.get_connection_from_secrets(azure_conn_id) # type: ignore[attr-defined] - azure_base_hook = azure_conn.get_hook() + azure_base_hook: AzureBaseHook = azure_conn.get_hook() scope = conf.get("postgres", "azure_oauth_scope", fallback=self.default_azure_oauth_scope) try: token = azure_base_hook.get_token(scope).token @@ -533,7 +536,7 @@ def get_azure_iam_token(self, conn: Connection) -> tuple[str, str, int]: "Please upgrade apache-airflow-providers-microsoft-azure>=12.8.0." ) from e raise - return cast("str", conn.login or azure_conn.login), cast("str", token), conn.port or 5432 + return cast("str", conn.login or azure_conn.login), token, conn.port or 5432 def get_table_primary_key(self, table: str, schema: str | None = "public") -> list[str] | None: """ From ed410c3d5bd4f921ac78d4927a48d88bb2296ad6 Mon Sep 17 00:00:00 2001 From: Karun Poudel Date: Sat, 4 Oct 2025 06:25:33 +0000 Subject: [PATCH 15/17] try --- dev/breeze/tests/test_selective_checks.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/dev/breeze/tests/test_selective_checks.py b/dev/breeze/tests/test_selective_checks.py index 5c00a026921b3..58c1ff8ce7a65 100644 --- a/dev/breeze/tests/test_selective_checks.py +++ b/dev/breeze/tests/test_selective_checks.py @@ -645,7 +645,7 @@ def assert_outputs_are_printed(expected_outputs: dict[str, str], stderr: str): ), { "selected-providers-list-as-string": "amazon common.sql google " - "openlineage pgvector postgres", + "microsoft.azure openlineage pgvector postgres", "all-python-versions": f"['{DEFAULT_PYTHON_MAJOR_MINOR_VERSION}']", "all-python-versions-list-as-string": DEFAULT_PYTHON_MAJOR_MINOR_VERSION, "python-versions": f"['{DEFAULT_PYTHON_MAJOR_MINOR_VERSION}']", From 8167b3e8131e255f394d0e918dbe5eeecdc6d4d6 Mon Sep 17 00:00:00 2001 From: Karun Poudel Date: Sat, 4 Oct 2025 06:32:34 +0000 Subject: [PATCH 16/17] pass --- dev/breeze/tests/test_selective_checks.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/dev/breeze/tests/test_selective_checks.py b/dev/breeze/tests/test_selective_checks.py index 58c1ff8ce7a65..a972f9d11ec14 100644 --- a/dev/breeze/tests/test_selective_checks.py +++ b/dev/breeze/tests/test_selective_checks.py @@ -667,7 +667,7 @@ def assert_outputs_are_printed(expected_outputs: dict[str, str], stderr: str): { "description": "amazon...google", "test_types": "Providers[amazon] " - "Providers[common.sql,openlineage,pgvector,postgres] " + "Providers[common.sql,microsoft.azure,openlineage,pgvector,postgres] " "Providers[google]", } ] From ad00901dd7843b84812f79d69e53aafac21d7786 Mon Sep 17 00:00:00 2001 From: Karun Poudel Date: Sat, 4 Oct 2025 18:53:36 +0000 Subject: [PATCH 17/17] doc and test --- .../postgres/docs/connections/postgres.rst | 2 +- .../providers/postgres/hooks/postgres.py | 7 ++++-- .../unit/postgres/hooks/test_postgres.py | 24 +++++++++++++++++++ 3 files changed, 30 insertions(+), 3 deletions(-) diff --git a/providers/postgres/docs/connections/postgres.rst b/providers/postgres/docs/connections/postgres.rst index 2583c199f6ed5..3018769afbca2 100644 --- a/providers/postgres/docs/connections/postgres.rst +++ b/providers/postgres/docs/connections/postgres.rst @@ -107,7 +107,7 @@ Extra (optional) for which you are requesting credentials. This parameter is case sensitive. If not specified than hostname from **Connection Host** is used. * ``azure_conn_id`` - Azure Connection ID to be used for authentication via Azure Entra ID. Azure Oauth token - is retrieved from the azure connection which is used as password for PostgreSQL connection. Requires `apache-airflow-providers-microsoft-azure>=12.8.0`. + is retrieved from the azure connection which is used as password for PostgreSQL connection. Scope for the Azure OAuth token can be set in the config option ``azure_oauth_scope`` under the section ``[postgres]``. Requires `apache-airflow-providers-microsoft-azure>=12.8.0`. Example "extras" field (Amazon RDS PostgreSQL or Amazon Aurora PostgreSQL): diff --git a/providers/postgres/src/airflow/providers/postgres/hooks/postgres.py b/providers/postgres/src/airflow/providers/postgres/hooks/postgres.py index 622c5960a926c..5fecc36556a20 100644 --- a/providers/postgres/src/airflow/providers/postgres/hooks/postgres.py +++ b/providers/postgres/src/airflow/providers/postgres/hooks/postgres.py @@ -516,6 +516,7 @@ def get_azure_iam_token(self, conn: Connection) -> tuple[str, str, int]: Get the Azure IAM token. This uses AzureBaseHook to retrieve an OAUTH token to connect to Postgres. + Scope for the OAuth token can be set in the config option ``azure_oauth_scope`` under the section ``[postgres]``. """ if TYPE_CHECKING: from airflow.providers.microsoft.azure.hooks.base_azure import AzureBaseHook @@ -530,10 +531,12 @@ def get_azure_iam_token(self, conn: Connection) -> tuple[str, str, int]: try: token = azure_base_hook.get_token(scope).token except AttributeError as e: - if "get_token" in str(e): + if e.name == "get_token" and e.obj == azure_base_hook: raise AttributeError( "'AzureBaseHook' object has no attribute 'get_token'. " - "Please upgrade apache-airflow-providers-microsoft-azure>=12.8.0." + "Please upgrade apache-airflow-providers-microsoft-azure>=12.8.0", + name=e.name, + obj=e.obj, ) from e raise return cast("str", conn.login or azure_conn.login), token, conn.port or 5432 diff --git a/providers/postgres/tests/unit/postgres/hooks/test_postgres.py b/providers/postgres/tests/unit/postgres/hooks/test_postgres.py index eaf38a7128079..47fffe0b6bb0b 100644 --- a/providers/postgres/tests/unit/postgres/hooks/test_postgres.py +++ b/providers/postgres/tests/unit/postgres/hooks/test_postgres.py @@ -471,6 +471,30 @@ def test_get_conn_azure_iam(self, mocker, mock_connect): assert mock_db_token in self.db_hook.sqlalchemy_url + def test_get_azure_iam_token_expect_failure_on_get_token(self, mocker): + """Test get_azure_iam_token method gets token from provided connection id""" + + class MockAzureBaseHookWithoutGetToken: + def __init__(self): + pass + + azure_conn_id = "azure_test_conn" + mock_connection_class = mocker.patch("airflow.providers.postgres.hooks.postgres.Connection") + mock_connection_class.get.return_value.get_hook.return_value = MockAzureBaseHookWithoutGetToken() + + self.connection.extra = json.dumps({"iam": True, "azure_conn_id": azure_conn_id}) + with pytest.raises( + AttributeError, + match=( + "'AzureBaseHook' object has no attribute 'get_token'. " + "Please upgrade apache-airflow-providers-microsoft-azure>=" + ), + ): + self.db_hook.get_azure_iam_token(self.connection) + + # Check AzureBaseHook initialization + mock_connection_class.get.assert_called_once_with(azure_conn_id) + def test_get_uri_from_connection_without_database_override(self, mocker): expected: str = f"postgresql{'+psycopg' if USE_PSYCOPG3 else ''}://login:password@host:1/database" self.db_hook.get_connection = mocker.MagicMock(