Skip to content
Closed
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
19 changes: 19 additions & 0 deletions providers/postgres/docs/configurations-ref.rst
Original file line number Diff line number Diff line change
@@ -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
15 changes: 14 additions & 1 deletion providers/postgres/docs/connections/postgres.rst
Original file line number Diff line number Diff line change
Expand Up @@ -96,14 +96,18 @@ Extra (optional)
* ``iam`` - If set to ``True`` than use AWS IAM database authentication for
`Amazon RDS <https://docs.aws.amazon.com/AmazonRDS/latest/UserGuide/UsingWithRDS.IAMDBAuth.html>`__,
`Amazon Aurora <https://docs.aws.amazon.com/AmazonRDS/latest/AuroraUserGuide/UsingWithRDS.IAMDBAuth.html>`__
or `Amazon Redshift <https://docs.aws.amazon.com/redshift/latest/mgmt/generating-user-credentials.html>`__.
`Amazon Redshift <https://docs.aws.amazon.com/redshift/latest/mgmt/generating-user-credentials.html>`__
or use Microsoft Entra Authentication for
`Azure Postgres Flexible Server <https://learn.microsoft.com/en-us/azure/postgresql/flexible-server/security-entra-concepts>`__.
* ``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.
If set to ``True`` than authenticate to Amazon Redshift Cluster, otherwise to Amazon RDS or Amazon Aurora.
* ``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):

Expand All @@ -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).
Expand Down
1 change: 1 addition & 0 deletions providers/postgres/docs/index.rst
Original file line number Diff line number Diff line change
Expand Up @@ -41,6 +41,7 @@
:maxdepth: 1
:caption: References

Configuration <configurations-ref>
Python API <_api/airflow/providers/postgres/index>
Dialects <dialects>

Expand Down
13 changes: 13 additions & 0 deletions providers/postgres/provider.yaml
Original file line number Diff line number Diff line change
Expand Up @@ -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"
Original file line number Diff line number Diff line change
Expand Up @@ -30,13 +30,19 @@
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,
)
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
Expand Down Expand Up @@ -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, ...]

Expand Down Expand Up @@ -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
Expand All @@ -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),
Expand Down Expand Up @@ -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
Expand Down Expand Up @@ -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.
Expand Down
27 changes: 27 additions & 0 deletions providers/postgres/tests/unit/postgres/hooks/test_postgres.py
Original file line number Diff line number Diff line change
Expand Up @@ -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(
Expand Down
Loading