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/connections/postgres.rst b/providers/postgres/docs/connections/postgres.rst index 539620ad08cc9..46d07991ec15a 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,8 @@ 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): @@ -125,6 +129,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/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 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" diff --git a/providers/postgres/src/airflow/providers/postgres/hooks/postgres.py b/providers/postgres/src/airflow/providers/postgres/hooks/postgres.py index 7a3be0ff4e314..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, @@ -37,6 +38,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 +70,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 +157,9 @@ class PostgresHook(DbApiHook): "aws_conn_id", "sqlalchemy_scheme", "sqlalchemy_query", + "azure_conn_id", } + 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 @@ -177,6 +180,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 +446,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 +511,19 @@ 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() + 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: """ 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..7a8f1acd71b87 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.default_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(