diff --git a/airflow/providers/postgres/operators/postgres.py b/airflow/providers/postgres/operators/postgres.py index 71f49ef7f8bc3..3e33c8ca0689d 100644 --- a/airflow/providers/postgres/operators/postgres.py +++ b/airflow/providers/postgres/operators/postgres.py @@ -18,7 +18,7 @@ from __future__ import annotations import warnings -from typing import Mapping, Sequence +from typing import Mapping from airflow.exceptions import AirflowProviderDeprecationWarning from airflow.providers.common.sql.operators.sql import SQLExecuteQueryOperator @@ -46,9 +46,7 @@ class PostgresOperator(SQLExecuteQueryOperator): Deprecated - use `hook_params={'options': '-c '}` instead. """ - template_fields: Sequence[str] = ("sql",) - template_fields_renderers = {"sql": "postgresql"} - template_ext: Sequence[str] = (".sql",) + template_fields_renderers = {**SQLExecuteQueryOperator.template_fields_renderers, "sql": "postgresql"} ui_color = "#ededed" def __init__( diff --git a/tests/providers/postgres/operators/test_postgres.py b/tests/providers/postgres/operators/test_postgres.py index 5bf8ee09360f6..03fc6fca54e99 100644 --- a/tests/providers/postgres/operators/test_postgres.py +++ b/tests/providers/postgres/operators/test_postgres.py @@ -190,3 +190,19 @@ def test_postgres_operator_openlineage_explicit_schema(self): assert lineage_on_complete.outputs[0].namespace == "postgres://postgres:5432" assert lineage_on_complete.outputs[0].name == "airflow.public.test_airflow" assert "schema" in lineage_on_complete.outputs[0].facets + + +def test_parameters_are_templatized(create_task_instance_of_operator): + """Test that PostgreSQL operator could template the same fields as SQLExecuteQueryOperator""" + ti = create_task_instance_of_operator( + PostgresOperator, + postgres_conn_id="{{ param.conn_id }}", + sql="SELECT * FROM {{ param.table }} WHERE spam = %(spam)s;", + parameters={"spam": "{{ param.bar }}"}, + dag_id="test-postgres-op-parameters-are-templatized", + task_id="test-task", + ) + task: PostgresOperator = ti.render_templates({"param": {"conn_id": "pg", "table": "foo", "bar": "egg"}}) + assert task.conn_id == "pg" + assert task.sql == "SELECT * FROM foo WHERE spam = %(spam)s;" + assert task.parameters == {"spam": "egg"}