diff --git a/providers/common/sql/src/airflow/providers/common/sql/dialects/dialect.py b/providers/common/sql/src/airflow/providers/common/sql/dialects/dialect.py index 0d4cc3d11e811..88453c73a33e5 100644 --- a/providers/common/sql/src/airflow/providers/common/sql/dialects/dialect.py +++ b/providers/common/sql/src/airflow/providers/common/sql/dialects/dialect.py @@ -33,7 +33,7 @@ class Dialect(LoggingMixin): """Generic dialect implementation.""" - pattern = re.compile(r'"([a-zA-Z0-9_]+)"') + pattern = re.compile(r"[^\w]") def __init__(self, hook, **kwargs) -> None: super().__init__(**kwargs) @@ -45,12 +45,6 @@ def __init__(self, hook, **kwargs) -> None: self.hook: DbApiHook = hook - @classmethod - def remove_quotes(cls, value: str | None) -> str | None: - if value: - return cls.pattern.sub(r"\1", value) - return value - @property def placeholder(self) -> str: return self.hook.placeholder @@ -60,16 +54,56 @@ def inspector(self) -> Inspector: return self.hook.inspector @property - def _insert_statement_format(self) -> str: - return self.hook._insert_statement_format # type: ignore + def insert_statement_format(self) -> str: + return self.hook.insert_statement_format + + @property + def replace_statement_format(self) -> str: + return self.hook.replace_statement_format @property - def _replace_statement_format(self) -> str: - return self.hook._replace_statement_format # type: ignore + def escape_word_format(self) -> str: + return self.hook.escape_word_format @property - def _escape_column_name_format(self) -> str: - return self.hook._escape_column_name_format # type: ignore + def escape_column_names(self) -> bool: + return self.hook.escape_column_names + + def escape_word(self, word: str) -> str: + """ + Escape the word if necessary. + + If the word is a reserved word or contains special characters or if the ``escape_column_names`` + property is set to True in connection extra field, then the given word will be escaped. + + :param word: Name of the column + :return: The escaped word + """ + if word != self.escape_word_format.format(self.unescape_word(word)) and ( + self.escape_column_names or word.casefold() in self.reserved_words or self.pattern.search(word) + ): + return self.escape_word_format.format(word) + return word + + def unescape_word(self, word: str | None) -> str | None: + """ + Remove escape characters from each part of a dotted identifier (e.g., schema.table). + + :param word: Escaped schema, table, or column name, potentially with multiple segments. + :return: The word without escaped characters. + """ + if not word: + return word + + escape_char_start = self.escape_word_format[0] + escape_char_end = self.escape_word_format[-1] + + def unescape_part(part: str) -> str: + if part.startswith(escape_char_start) and part.endswith(escape_char_end): + return part[1:-1] + return part + + return ".".join(map(unescape_part, word.split("."))) @classmethod def extract_schema_from_table(cls, table: str) -> tuple[str, str | None]: @@ -87,8 +121,8 @@ def get_column_names( for column in filter( predicate, self.inspector.get_columns( - table_name=self.remove_quotes(table), - schema=self.remove_quotes(schema) if schema else None, + table_name=self.unescape_word(table), + schema=self.unescape_word(schema) if schema else None, ), ) ) @@ -110,8 +144,8 @@ def get_primary_keys(self, table: str, schema: str | None = None) -> list[str] | if schema is None: table, schema = self.extract_schema_from_table(table) primary_keys = self.inspector.get_pk_constraint( - table_name=self.remove_quotes(table), - schema=self.remove_quotes(schema) if schema else None, + table_name=self.unescape_word(table), + schema=self.unescape_word(schema) if schema else None, ).get("constrained_columns", []) self.log.debug("Primary keys for table '%s': %s", table, primary_keys) return primary_keys @@ -138,20 +172,6 @@ def get_records( def reserved_words(self) -> set[str]: return self.hook.reserved_words - def escape_column_name(self, column_name: str) -> str: - """ - Escape the column name if it's a reserved word. - - :param column_name: Name of the column - :return: The escaped column name if needed - """ - if ( - column_name != self._escape_column_name_format.format(column_name) - and column_name.casefold() in self.reserved_words - ): - return self._escape_column_name_format.format(column_name) - return column_name - def _joined_placeholders(self, values) -> str: placeholders = [ self.placeholder, @@ -160,7 +180,7 @@ def _joined_placeholders(self, values) -> str: def _joined_target_fields(self, target_fields) -> str: if target_fields: - target_fields = ", ".join(map(self.escape_column_name, target_fields)) + target_fields = ", ".join(map(self.escape_word, target_fields)) return f"({target_fields})" return "" @@ -173,7 +193,7 @@ def generate_insert_sql(self, table, values, target_fields, **kwargs) -> str: :param target_fields: The names of the columns to fill in the table :return: The generated INSERT SQL statement """ - return self._insert_statement_format.format( + return self.insert_statement_format.format( table, self._joined_target_fields(target_fields), self._joined_placeholders(values) ) @@ -186,6 +206,6 @@ def generate_replace_sql(self, table, values, target_fields, **kwargs) -> str: :param target_fields: The names of the columns to fill in the table :return: The generated REPLACE SQL statement """ - return self._replace_statement_format.format( + return self.replace_statement_format.format( table, self._joined_target_fields(target_fields), self._joined_placeholders(values) ) diff --git a/providers/common/sql/src/airflow/providers/common/sql/dialects/dialect.pyi b/providers/common/sql/src/airflow/providers/common/sql/dialects/dialect.pyi index 4c71747f19f1e..ff208e77e5c82 100644 --- a/providers/common/sql/src/airflow/providers/common/sql/dialects/dialect.pyi +++ b/providers/common/sql/src/airflow/providers/common/sql/dialects/dialect.pyi @@ -45,11 +45,17 @@ T = TypeVar("T") class Dialect(LoggingMixin): hook: Incomplete def __init__(self, hook, **kwargs) -> None: ... - @classmethod - def remove_quotes(cls, value: str | None) -> str | None: ... + def escape_word(self, column_name: str) -> str: ... + def unescape_word(self, value: str | None) -> str | None: ... @property def placeholder(self) -> str: ... @property + def insert_statement_format(self) -> str: ... + @property + def replace_statement_format(self) -> str: ... + @property + def escape_word_format(self) -> str: ... + @property def inspector(self) -> Inspector: ... @classmethod def extract_schema_from_table(cls, table: str) -> tuple[str, str | None]: ... @@ -72,6 +78,5 @@ class Dialect(LoggingMixin): ) -> Any: ... @property def reserved_words(self) -> set[str]: ... - def escape_column_name(self, column_name: str) -> str: ... def generate_insert_sql(self, table, values, target_fields, **kwargs) -> str: ... def generate_replace_sql(self, table, values, target_fields, **kwargs) -> str: ... diff --git a/providers/common/sql/src/airflow/providers/common/sql/hooks/sql.py b/providers/common/sql/src/airflow/providers/common/sql/hooks/sql.py index fec2b81ec128d..ff4fd2843e786 100644 --- a/providers/common/sql/src/airflow/providers/common/sql/hooks/sql.py +++ b/providers/common/sql/src/airflow/providers/common/sql/hooks/sql.py @@ -47,6 +47,7 @@ ) from airflow.hooks.base import BaseHook from airflow.providers.common.sql.dialects.dialect import Dialect +from airflow.providers.common.sql.hooks import handlers from airflow.utils.module_loading import import_string if TYPE_CHECKING: @@ -67,24 +68,18 @@ def return_single_query_results(sql: str | Iterable[str], return_last: bool, split_statements: bool | None): warnings.warn(WARNING_MESSAGE.format("return_single_query_results"), DeprecationWarning, stacklevel=2) - from airflow.providers.common.sql.hooks import handlers - return handlers.return_single_query_results(sql, return_last, split_statements) def fetch_all_handler(cursor) -> list[tuple] | None: warnings.warn(WARNING_MESSAGE.format("fetch_all_handler"), DeprecationWarning, stacklevel=2) - from airflow.providers.common.sql.hooks import handlers - return handlers.fetch_all_handler(cursor) def fetch_one_handler(cursor) -> list[tuple] | None: warnings.warn(WARNING_MESSAGE.format("fetch_one_handler"), DeprecationWarning, stacklevel=2) - from airflow.providers.common.sql.hooks import handlers - return handlers.fetch_one_handler(cursor) @@ -184,13 +179,10 @@ def __init__(self, *args, schema: str | None = None, log_sql: bool = True, **kwa self.__schema = schema self.log_sql = log_sql self.descriptions: list[Sequence[Sequence] | None] = [] - self._insert_statement_format: str = kwargs.get( - "insert_statement_format", "INSERT INTO {} {} VALUES ({})" - ) - self._replace_statement_format: str = kwargs.get( - "replace_statement_format", "REPLACE INTO {} {} VALUES ({})" - ) - self._escape_column_name_format: str = kwargs.get("escape_column_name_format", '"{}"') + self._insert_statement_format: str | None = kwargs.get("insert_statement_format") + self._replace_statement_format: str | None = kwargs.get("replace_statement_format") + self._escape_word_format: str | None = kwargs.get("escape_word_format") + self._escape_column_names: bool | None = kwargs.get("escape_column_names") self._connection: Connection | None = kwargs.pop("connection", None) def get_conn_id(self) -> str: @@ -212,6 +204,38 @@ def placeholder(self) -> str: ) return self._placeholder + @property + def insert_statement_format(self) -> str: + """Return the insert statement format.""" + if not self._insert_statement_format: + self._insert_statement_format = self.connection_extra.get( + "insert_statement_format", "INSERT INTO {} {} VALUES ({})" + ) + return self._insert_statement_format + + @property + def replace_statement_format(self) -> str: + """Return the replacement statement format.""" + if not self._replace_statement_format: + self._replace_statement_format = self.connection_extra.get( + "replace_statement_format", "REPLACE INTO {} {} VALUES ({})" + ) + return self._replace_statement_format + + @property + def escape_word_format(self) -> str: + """Return the escape word format.""" + if not self._escape_word_format: + self._escape_word_format = self.connection_extra.get("escape_word_format", '"{}"') + return self._escape_word_format + + @property + def escape_column_names(self) -> bool: + """Return the escape column names flag.""" + if not self._escape_column_names: + self._escape_column_names = self.connection_extra.get("escape_column_names", False) + return self._escape_column_names + @property def connection(self) -> Connection: if self._connection is None: @@ -413,7 +437,7 @@ def get_records( :param sql: the sql statement to be executed (str) or a list of sql statements to execute :param parameters: The parameters to render the SQL query with. """ - return self.run(sql=sql, parameters=parameters, handler=fetch_all_handler) + return self.run(sql=sql, parameters=parameters, handler=handlers.fetch_all_handler) def get_first(self, sql: str | list[str], parameters: Iterable | Mapping[str, Any] | None = None) -> Any: """ @@ -422,7 +446,7 @@ def get_first(self, sql: str | list[str], parameters: Iterable | Mapping[str, An :param sql: the sql statement to be executed (str) or a list of sql statements to execute :param parameters: The parameters to render the SQL query with. """ - return self.run(sql=sql, parameters=parameters, handler=fetch_one_handler) + return self.run(sql=sql, parameters=parameters, handler=handlers.fetch_one_handler) @staticmethod def strip_sql_string(sql: str) -> str: @@ -557,7 +581,7 @@ def run( if handler is not None: result = self._make_common_data_structure(handler(cur)) - if return_single_query_results(sql, return_last, split_statements): + if handlers.return_single_query_results(sql, return_last, split_statements): _last_result = result _last_description = cur.description else: @@ -572,7 +596,7 @@ def run( if handler is None: return None - if return_single_query_results(sql, return_last, split_statements): + if handlers.return_single_query_results(sql, return_last, split_statements): self.descriptions = [_last_description] return _last_result else: diff --git a/providers/common/sql/src/airflow/providers/common/sql/hooks/sql.pyi b/providers/common/sql/src/airflow/providers/common/sql/hooks/sql.pyi index 23f8ee6c17ca8..fb9272e229091 100644 --- a/providers/common/sql/src/airflow/providers/common/sql/hooks/sql.pyi +++ b/providers/common/sql/src/airflow/providers/common/sql/hooks/sql.pyi @@ -72,6 +72,14 @@ class DbApiHook(BaseHook): @cached_property def placeholder(self) -> str: ... @property + def insert_statement_format(self) -> str: ... + @property + def replace_statement_format(self) -> str: ... + @property + def escape_word_format(self) -> str: ... + @property + def escape_column_names(self) -> bool: ... + @property def connection(self) -> Connection: ... @connection.setter def connection(self, value: Any) -> None: ... diff --git a/providers/common/sql/tests/provider_tests/common/sql/dialects/test_dialect.py b/providers/common/sql/tests/provider_tests/common/sql/dialects/test_dialect.py index 1021b1c617caa..e007969befa65 100644 --- a/providers/common/sql/tests/provider_tests/common/sql/dialects/test_dialect.py +++ b/providers/common/sql/tests/provider_tests/common/sql/dialects/test_dialect.py @@ -29,18 +29,67 @@ class TestDialect: def setup_method(self): inspector = MagicMock(spc=Inspector) inspector.get_columns.side_effect = lambda table_name, schema: [ - {"name": "id", "identity": True}, + {"name": "index", "identity": True}, {"name": "name"}, {"name": "firstname"}, {"name": "age"}, ] inspector.get_pk_constraint.side_effect = lambda table_name, schema: {"constrained_columns": ["id"]} self.test_db_hook = MagicMock(placeholder="?", inspector=inspector, spec=DbApiHook) + self.test_db_hook.reserved_words = {"index", "user"} + self.test_db_hook.insert_statement_format = "INSERT INTO {} {} VALUES ({})" + self.test_db_hook.replace_statement_format = "REPLACE INTO {} {} VALUES ({})" + self.test_db_hook.escape_word_format = '"{}"' + self.test_db_hook.escape_column_names = False - def test_remove_quotes(self): - assert not Dialect.remove_quotes(None) - assert Dialect.remove_quotes("table") == "table" - assert Dialect.remove_quotes('"table"') == "table" + def test_insert_statement_format(self): + assert Dialect(self.test_db_hook).insert_statement_format == "INSERT INTO {} {} VALUES ({})" + + def test_replace_statement_format(self): + assert Dialect(self.test_db_hook).replace_statement_format == "REPLACE INTO {} {} VALUES ({})" + + def test_escape_word_format(self): + assert Dialect(self.test_db_hook).escape_word_format == '"{}"' + + def test_unescape_word(self): + assert Dialect(self.test_db_hook).unescape_word('"table"') == "table" + + def test_unescape_word_with_different_format(self): + self.test_db_hook.escape_word_format = "[{}]" + dialect = Dialect(self.test_db_hook) + assert not dialect.unescape_word(None) + assert dialect.unescape_word("table") == "table" + assert dialect.unescape_word("t@ble") == "t@ble" + assert dialect.unescape_word("table_name") == "table_name" + assert dialect.unescape_word('"table"') == '"table"' + assert dialect.unescape_word("[table]") == "table" + assert dialect.unescape_word("schema.[t@ble]") == "schema.t@ble" + assert dialect.unescape_word("[schema].[t@ble]") == "schema.t@ble" + assert dialect.unescape_word("[schema].table") == "schema.table" + + def test_escape_word(self): + assert Dialect(self.test_db_hook).escape_word('"table"') == '"table"' + + def test_escape_word_with_different_format(self): + self.test_db_hook.escape_word_format = "[{}]" + dialect = Dialect(self.test_db_hook) + assert dialect.escape_word("name") == "name" + assert dialect.escape_word("[name]") == "[name]" + assert dialect.escape_word("n@me") == "[n@me]" + assert dialect.escape_word("index") == "[index]" + assert dialect.escape_word("User") == "[User]" + assert dialect.escape_word("attributes.id") == "[attributes.id]" + + def test_escape_word_when_all_column_names_must_be_escaped(self): + self.test_db_hook.escape_word_format = "[{}]" + self.test_db_hook.escape_column_names = True + dialect = Dialect(self.test_db_hook) + assert dialect.escape_word("name") == "[name]" + assert dialect.escape_word("[name]") == "[name]" + assert dialect.escape_word("n@me") == "[n@me]" + assert dialect.escape_word("index") == "[index]" + assert dialect.escape_word("User") == "[User]" + assert dialect.escape_word("attributes.id") == "[attributes.id]" def test_placeholder(self): assert Dialect(self.test_db_hook).placeholder == "?" @@ -50,7 +99,7 @@ def test_extract_schema_from_table(self): def test_get_column_names(self): assert Dialect(self.test_db_hook).get_column_names("table", "schema") == [ - "id", + "index", "name", "firstname", "age", @@ -65,3 +114,38 @@ def test_get_target_fields(self): def test_get_primary_keys(self): assert Dialect(self.test_db_hook).get_primary_keys("table", "schema") == ["id"] + + def test_generate_replace_sql(self): + values = [ + {"index": 1, "name": "Stallone", "firstname": "Sylvester", "age": "78"}, + {"index": 2, "name": "Statham", "firstname": "Jason", "age": "57"}, + {"index": 3, "name": "Li", "firstname": "Jet", "age": "61"}, + {"index": 4, "name": "Lundgren", "firstname": "Dolph", "age": "66"}, + {"index": 5, "name": "Norris", "firstname": "Chuck", "age": "84"}, + ] + target_fields = ["index", "name", "firstname", "age"] + sql = Dialect(self.test_db_hook).generate_replace_sql("hollywood.actors", values, target_fields) + assert ( + sql + == """ + REPLACE INTO hollywood.actors ("index", name, firstname, age) VALUES (?,?,?,?,?) + """.strip() + ) + + def test_generate_replace_sql_when_escape_column_names_is_enabled(self): + values = [ + {"index": 1, "name": "Stallone", "firstname": "Sylvester", "age": "78"}, + {"index": 2, "name": "Statham", "firstname": "Jason", "age": "57"}, + {"index": 3, "name": "Li", "firstname": "Jet", "age": "61"}, + {"index": 4, "name": "Lundgren", "firstname": "Dolph", "age": "66"}, + {"index": 5, "name": "Norris", "firstname": "Chuck", "age": "84"}, + ] + target_fields = ["index", "name", "firstname", "age"] + self.test_db_hook.escape_column_names = True + sql = Dialect(self.test_db_hook).generate_replace_sql("hollywood.actors", values, target_fields) + assert ( + sql + == """ + REPLACE INTO hollywood.actors ("index", "name", "firstname", "age") VALUES (?,?,?,?,?) + """.strip() + ) diff --git a/providers/common/sql/tests/provider_tests/common/sql/hooks/test_sql.py b/providers/common/sql/tests/provider_tests/common/sql/hooks/test_sql.py index 3588a15c9e2a9..21f0120b61af6 100644 --- a/providers/common/sql/tests/provider_tests/common/sql/hooks/test_sql.py +++ b/providers/common/sql/tests/provider_tests/common/sql/hooks/test_sql.py @@ -263,6 +263,11 @@ def test_placeholder_multiple_times_and_make_sure_connection_is_only_invoked_onc assert dbapi_hook.placeholder == "%s" assert dbapi_hook.connection_invocations == 1 + @pytest.mark.db_test + def test_escape_column_names(self): + dbapi_hook = mock_db_hook(DbApiHook) + assert not dbapi_hook.escape_column_names + @pytest.mark.db_test def test_dialect_name(self): dbapi_hook = mock_db_hook(DbApiHook) diff --git a/providers/src/airflow/providers/microsoft/mssql/dialects/mssql.py b/providers/src/airflow/providers/microsoft/mssql/dialects/mssql.py index fc2110a762d64..edad1a11515d5 100644 --- a/providers/src/airflow/providers/microsoft/mssql/dialects/mssql.py +++ b/providers/src/airflow/providers/microsoft/mssql/dialects/mssql.py @@ -47,7 +47,7 @@ def get_primary_keys(self, table: str, schema: str | None = None) -> list[str] | def generate_replace_sql(self, table, values, target_fields, **kwargs) -> str: primary_keys = self.get_primary_keys(table) columns = [ - self.escape_column_name(target_field) + self.escape_word(target_field) for target_field in target_fields if target_field in set(target_fields).difference(set(primary_keys)) ] @@ -56,9 +56,9 @@ def generate_replace_sql(self, table, values, target_fields, **kwargs) -> str: self.log.debug("columns: %s", columns) return f"""MERGE INTO {table} WITH (ROWLOCK) AS target - USING (SELECT {', '.join(map(lambda column: f'{self.placeholder} AS {column}', target_fields))}) AS source - ON {' AND '.join(map(lambda column: f'target.{self.escape_column_name(column)} = source.{column}', primary_keys))} + USING (SELECT {', '.join(map(lambda column: f'{self.placeholder} AS {self.escape_word(column)}', target_fields))}) AS source + ON {' AND '.join(map(lambda column: f'target.{self.escape_word(column)} = source.{self.escape_word(column)}', primary_keys))} WHEN MATCHED THEN UPDATE SET {', '.join(map(lambda column: f'target.{column} = source.{column}', columns))} WHEN NOT MATCHED THEN - INSERT ({', '.join(target_fields)}) VALUES ({', '.join(map(lambda column: f'source.{self.escape_column_name(column)}', target_fields))});""" + INSERT ({', '.join(map(self.escape_word, target_fields))}) VALUES ({', '.join(map(lambda column: f'source.{self.escape_word(column)}', target_fields))});""" diff --git a/providers/src/airflow/providers/microsoft/mssql/hooks/mssql.py b/providers/src/airflow/providers/microsoft/mssql/hooks/mssql.py index a29018ec0ff91..73a646a7f139f 100644 --- a/providers/src/airflow/providers/microsoft/mssql/hooks/mssql.py +++ b/providers/src/airflow/providers/microsoft/mssql/hooks/mssql.py @@ -55,7 +55,7 @@ def __init__( sqlalchemy_scheme: str | None = None, **kwargs, ) -> None: - super().__init__(*args, **kwargs) + super().__init__(*args, **{**kwargs, **{"escape_word_format": "[{}]"}}) self.schema = kwargs.pop("schema", None) self._sqlalchemy_scheme = sqlalchemy_scheme diff --git a/providers/src/airflow/providers/mysql/hooks/mysql.py b/providers/src/airflow/providers/mysql/hooks/mysql.py index 48185c1cf5516..05c851e7c5d76 100644 --- a/providers/src/airflow/providers/mysql/hooks/mysql.py +++ b/providers/src/airflow/providers/mysql/hooks/mysql.py @@ -78,11 +78,10 @@ class MySqlHook(DbApiHook): supports_autocommit = True def __init__(self, *args, **kwargs) -> None: - super().__init__(*args, **kwargs) + super().__init__(*args, **{**kwargs, **{"escape_word_format": "`{}`"}}) self.schema = kwargs.pop("schema", None) self.local_infile = kwargs.pop("local_infile", False) self.init_command = kwargs.pop("init_command", None) - self._escape_column_name_format: str = kwargs.get("escape_column_name_format", "`{}`") def set_autocommit(self, conn: MySQLConnectionTypes, autocommit: bool) -> None: """ diff --git a/providers/src/airflow/providers/postgres/dialects/postgres.py b/providers/src/airflow/providers/postgres/dialects/postgres.py index 5db4cca18f8e5..e8bef1d2b5a48 100644 --- a/providers/src/airflow/providers/postgres/dialects/postgres.py +++ b/providers/src/airflow/providers/postgres/dialects/postgres.py @@ -51,7 +51,7 @@ def get_primary_keys(self, table: str, schema: str | None = None) -> list[str] | and kcu.table_name = %s """ pk_columns = [ - row[0] for row in self.get_records(sql, (self.remove_quotes(schema), self.remove_quotes(table))) + row[0] for row in self.get_records(sql, (self.unescape_word(schema), self.unescape_word(table))) ] return pk_columns or None @@ -79,8 +79,8 @@ def generate_replace_sql(self, table, values, target_fields, **kwargs) -> str: replace_index = [replace_index] sql = self.generate_insert_sql(table, values, target_fields, **kwargs) - on_conflict_str = f" ON CONFLICT ({', '.join(map(self.escape_column_name, replace_index))})" - replace_target = [self.escape_column_name(f) for f in target_fields if f not in replace_index] + on_conflict_str = f" ON CONFLICT ({', '.join(map(self.escape_word, replace_index))})" + replace_target = [self.escape_word(f) for f in target_fields if f not in replace_index] if replace_target: replace_target_str = ", ".join(f"{col} = excluded.{col}" for col in replace_target) diff --git a/providers/tests/microsoft/mssql/dialects/test_mssql.py b/providers/tests/microsoft/mssql/dialects/test_mssql.py index 762c4c463dfdb..749a79c13fcd1 100644 --- a/providers/tests/microsoft/mssql/dialects/test_mssql.py +++ b/providers/tests/microsoft/mssql/dialects/test_mssql.py @@ -29,21 +29,23 @@ class TestMsSqlDialect: def setup_method(self): inspector = MagicMock(spc=Inspector) inspector.get_columns.side_effect = lambda table_name, schema: [ - {"name": "id", "identity": True}, + {"name": "index", "identity": True}, {"name": "name"}, {"name": "firstname"}, {"name": "age"}, ] self.test_db_hook = MagicMock(placeholder="?", inspector=inspector, spec=DbApiHook) - self.test_db_hook.run.side_effect = lambda *args: [("id",)] - self.test_db_hook._escape_column_name_format = '"{}"' + self.test_db_hook.run.side_effect = lambda *args: [("index",)] + self.test_db_hook.reserved_words = {"index", "user"} + self.test_db_hook.escape_word_format = "[{}]" + self.test_db_hook.escape_column_names = False def test_placeholder(self): assert MsSqlDialect(self.test_db_hook).placeholder == "?" def test_get_column_names(self): assert MsSqlDialect(self.test_db_hook).get_column_names("hollywood.actors") == [ - "id", + "index", "name", "firstname", "age", @@ -57,27 +59,51 @@ def test_get_target_fields(self): ] def test_get_primary_keys(self): - assert MsSqlDialect(self.test_db_hook).get_primary_keys("hollywood.actors") == ["id"] + assert MsSqlDialect(self.test_db_hook).get_primary_keys("hollywood.actors") == ["index"] def test_generate_replace_sql(self): values = [ - {"id": "id", "name": "Stallone", "firstname": "Sylvester", "age": "78"}, - {"id": "id", "name": "Statham", "firstname": "Jason", "age": "57"}, - {"id": "id", "name": "Li", "firstname": "Jet", "age": "61"}, - {"id": "id", "name": "Lundgren", "firstname": "Dolph", "age": "66"}, - {"id": "id", "name": "Norris", "firstname": "Chuck", "age": "84"}, + {"index": 1, "name": "Stallone", "firstname": "Sylvester", "age": "78"}, + {"index": 2, "name": "Statham", "firstname": "Jason", "age": "57"}, + {"index": 3, "name": "Li", "firstname": "Jet", "age": "61"}, + {"index": 4, "name": "Lundgren", "firstname": "Dolph", "age": "66"}, + {"index": 5, "name": "Norris", "firstname": "Chuck", "age": "84"}, ] - target_fields = ["id", "name", "firstname", "age"] + target_fields = ["index", "name", "firstname", "age"] sql = MsSqlDialect(self.test_db_hook).generate_replace_sql("hollywood.actors", values, target_fields) assert ( sql == """ MERGE INTO hollywood.actors WITH (ROWLOCK) AS target - USING (SELECT ? AS id, ? AS name, ? AS firstname, ? AS age) AS source - ON target.id = source.id + USING (SELECT ? AS [index], ? AS name, ? AS firstname, ? AS age) AS source + ON target.[index] = source.[index] WHEN MATCHED THEN UPDATE SET target.name = source.name, target.firstname = source.firstname, target.age = source.age WHEN NOT MATCHED THEN - INSERT (id, name, firstname, age) VALUES (source.id, source.name, source.firstname, source.age); + INSERT ([index], name, firstname, age) VALUES (source.[index], source.name, source.firstname, source.age); + """.strip() + ) + + def test_generate_replace_sql_when_escape_column_names_is_enabled(self): + values = [ + {"index": 1, "name": "Stallone", "firstname": "Sylvester", "age": "78"}, + {"index": 2, "name": "Statham", "firstname": "Jason", "age": "57"}, + {"index": 3, "name": "Li", "firstname": "Jet", "age": "61"}, + {"index": 4, "name": "Lundgren", "firstname": "Dolph", "age": "66"}, + {"index": 5, "name": "Norris", "firstname": "Chuck", "age": "84"}, + ] + target_fields = ["index", "name", "firstname", "age"] + self.test_db_hook.escape_column_names = True + sql = MsSqlDialect(self.test_db_hook).generate_replace_sql("hollywood.actors", values, target_fields) + assert ( + sql + == """ + MERGE INTO hollywood.actors WITH (ROWLOCK) AS target + USING (SELECT ? AS [index], ? AS [name], ? AS [firstname], ? AS [age]) AS source + ON target.[index] = source.[index] + WHEN MATCHED THEN + UPDATE SET target.[name] = source.[name], target.[firstname] = source.[firstname], target.[age] = source.[age] + WHEN NOT MATCHED THEN + INSERT ([index], [name], [firstname], [age]) VALUES (source.[index], source.[name], source.[firstname], source.[age]); """.strip() ) diff --git a/providers/tests/microsoft/mssql/hooks/test_mssql.py b/providers/tests/microsoft/mssql/hooks/test_mssql.py index 7153edde85217..8f5da16722a0a 100644 --- a/providers/tests/microsoft/mssql/hooks/test_mssql.py +++ b/providers/tests/microsoft/mssql/hooks/test_mssql.py @@ -246,7 +246,7 @@ def test_get_sqlalchemy_engine(self, get_connection, mssql_connections): def test_generate_insert_sql(self, get_connection): get_connection.return_value = PYMSSQL_CONN - hook = MsSqlHook() + hook = MsSqlHook(escape_word_format="[{}]") sql = hook._generate_insert_sql( table="YAMMER_GROUPS_ACTIVITY_DETAIL", values=[ diff --git a/providers/tests/microsoft/mssql/resources/replace.sql b/providers/tests/microsoft/mssql/resources/replace.sql index 07c7ec29e0188..ee44d857327da 100644 --- a/providers/tests/microsoft/mssql/resources/replace.sql +++ b/providers/tests/microsoft/mssql/resources/replace.sql @@ -18,9 +18,9 @@ */ MERGE INTO YAMMER_GROUPS_ACTIVITY_DETAIL WITH (ROWLOCK) AS target - USING (SELECT %s AS ReportRefreshDate, %s AS UserId, %s AS UserPrincipalName, %s AS LastActivityDate, %s AS IsDeleted, %s AS DeletedDate, %s AS AssignedProducts, %s AS TeamChatMessageCount, %s AS PrivateChatMessageCount, %s AS CallCount, %s AS MeetingCount, %s AS MeetingsOrganizedCount, %s AS MeetingsAttendedCount, %s AS AdHocMeetingsOrganizedCount, %s AS AdHocMeetingsAttendedCount, %s AS ScheduledOne-timeMeetingsOrganizedCount, %s AS ScheduledOne-timeMeetingsAttendedCount, %s AS ScheduledRecurringMeetingsOrganizedCount, %s AS ScheduledRecurringMeetingsAttendedCount, %s AS AudioDuration, %s AS VideoDuration, %s AS ScreenShareDuration, %s AS AudioDurationInSeconds, %s AS VideoDurationInSeconds, %s AS ScreenShareDurationInSeconds, %s AS HasOtherAction, %s AS UrgentMessages, %s AS PostMessages, %s AS TenantDisplayName, %s AS SharedChannelTenantDisplayNames, %s AS ReplyMessages, %s AS IsLicensed, %s AS ReportPeriod, %s AS LoadDate) AS source + USING (SELECT %s AS ReportRefreshDate, %s AS UserId, %s AS UserPrincipalName, %s AS LastActivityDate, %s AS IsDeleted, %s AS DeletedDate, %s AS AssignedProducts, %s AS TeamChatMessageCount, %s AS PrivateChatMessageCount, %s AS CallCount, %s AS MeetingCount, %s AS MeetingsOrganizedCount, %s AS MeetingsAttendedCount, %s AS AdHocMeetingsOrganizedCount, %s AS AdHocMeetingsAttendedCount, %s AS [ScheduledOne-timeMeetingsOrganizedCount], %s AS [ScheduledOne-timeMeetingsAttendedCount], %s AS ScheduledRecurringMeetingsOrganizedCount, %s AS ScheduledRecurringMeetingsAttendedCount, %s AS AudioDuration, %s AS VideoDuration, %s AS ScreenShareDuration, %s AS AudioDurationInSeconds, %s AS VideoDurationInSeconds, %s AS ScreenShareDurationInSeconds, %s AS HasOtherAction, %s AS UrgentMessages, %s AS PostMessages, %s AS TenantDisplayName, %s AS SharedChannelTenantDisplayNames, %s AS ReplyMessages, %s AS IsLicensed, %s AS ReportPeriod, %s AS LoadDate) AS source ON target.GroupDisplayName = source.GroupDisplayName AND target.OwnerPrincipalName = source.OwnerPrincipalName AND target.ReportPeriod = source.ReportPeriod AND target.ReportRefreshDate = source.ReportRefreshDate WHEN MATCHED THEN - UPDATE SET target.UserId = source.UserId, target.UserPrincipalName = source.UserPrincipalName, target.LastActivityDate = source.LastActivityDate, target.IsDeleted = source.IsDeleted, target.DeletedDate = source.DeletedDate, target.AssignedProducts = source.AssignedProducts, target.TeamChatMessageCount = source.TeamChatMessageCount, target.PrivateChatMessageCount = source.PrivateChatMessageCount, target.CallCount = source.CallCount, target.MeetingCount = source.MeetingCount, target.MeetingsOrganizedCount = source.MeetingsOrganizedCount, target.MeetingsAttendedCount = source.MeetingsAttendedCount, target.AdHocMeetingsOrganizedCount = source.AdHocMeetingsOrganizedCount, target.AdHocMeetingsAttendedCount = source.AdHocMeetingsAttendedCount, target.ScheduledOne-timeMeetingsOrganizedCount = source.ScheduledOne-timeMeetingsOrganizedCount, target.ScheduledOne-timeMeetingsAttendedCount = source.ScheduledOne-timeMeetingsAttendedCount, target.ScheduledRecurringMeetingsOrganizedCount = source.ScheduledRecurringMeetingsOrganizedCount, target.ScheduledRecurringMeetingsAttendedCount = source.ScheduledRecurringMeetingsAttendedCount, target.AudioDuration = source.AudioDuration, target.VideoDuration = source.VideoDuration, target.ScreenShareDuration = source.ScreenShareDuration, target.AudioDurationInSeconds = source.AudioDurationInSeconds, target.VideoDurationInSeconds = source.VideoDurationInSeconds, target.ScreenShareDurationInSeconds = source.ScreenShareDurationInSeconds, target.HasOtherAction = source.HasOtherAction, target.UrgentMessages = source.UrgentMessages, target.PostMessages = source.PostMessages, target.TenantDisplayName = source.TenantDisplayName, target.SharedChannelTenantDisplayNames = source.SharedChannelTenantDisplayNames, target.ReplyMessages = source.ReplyMessages, target.IsLicensed = source.IsLicensed, target.LoadDate = source.LoadDate + UPDATE SET target.UserId = source.UserId, target.UserPrincipalName = source.UserPrincipalName, target.LastActivityDate = source.LastActivityDate, target.IsDeleted = source.IsDeleted, target.DeletedDate = source.DeletedDate, target.AssignedProducts = source.AssignedProducts, target.TeamChatMessageCount = source.TeamChatMessageCount, target.PrivateChatMessageCount = source.PrivateChatMessageCount, target.CallCount = source.CallCount, target.MeetingCount = source.MeetingCount, target.MeetingsOrganizedCount = source.MeetingsOrganizedCount, target.MeetingsAttendedCount = source.MeetingsAttendedCount, target.AdHocMeetingsOrganizedCount = source.AdHocMeetingsOrganizedCount, target.AdHocMeetingsAttendedCount = source.AdHocMeetingsAttendedCount, target.[ScheduledOne-timeMeetingsOrganizedCount] = source.[ScheduledOne-timeMeetingsOrganizedCount], target.[ScheduledOne-timeMeetingsAttendedCount] = source.[ScheduledOne-timeMeetingsAttendedCount], target.ScheduledRecurringMeetingsOrganizedCount = source.ScheduledRecurringMeetingsOrganizedCount, target.ScheduledRecurringMeetingsAttendedCount = source.ScheduledRecurringMeetingsAttendedCount, target.AudioDuration = source.AudioDuration, target.VideoDuration = source.VideoDuration, target.ScreenShareDuration = source.ScreenShareDuration, target.AudioDurationInSeconds = source.AudioDurationInSeconds, target.VideoDurationInSeconds = source.VideoDurationInSeconds, target.ScreenShareDurationInSeconds = source.ScreenShareDurationInSeconds, target.HasOtherAction = source.HasOtherAction, target.UrgentMessages = source.UrgentMessages, target.PostMessages = source.PostMessages, target.TenantDisplayName = source.TenantDisplayName, target.SharedChannelTenantDisplayNames = source.SharedChannelTenantDisplayNames, target.ReplyMessages = source.ReplyMessages, target.IsLicensed = source.IsLicensed, target.LoadDate = source.LoadDate WHEN NOT MATCHED THEN - INSERT (ReportRefreshDate, UserId, UserPrincipalName, LastActivityDate, IsDeleted, DeletedDate, AssignedProducts, TeamChatMessageCount, PrivateChatMessageCount, CallCount, MeetingCount, MeetingsOrganizedCount, MeetingsAttendedCount, AdHocMeetingsOrganizedCount, AdHocMeetingsAttendedCount, ScheduledOne-timeMeetingsOrganizedCount, ScheduledOne-timeMeetingsAttendedCount, ScheduledRecurringMeetingsOrganizedCount, ScheduledRecurringMeetingsAttendedCount, AudioDuration, VideoDuration, ScreenShareDuration, AudioDurationInSeconds, VideoDurationInSeconds, ScreenShareDurationInSeconds, HasOtherAction, UrgentMessages, PostMessages, TenantDisplayName, SharedChannelTenantDisplayNames, ReplyMessages, IsLicensed, ReportPeriod, LoadDate) VALUES (source.ReportRefreshDate, source.UserId, source.UserPrincipalName, source.LastActivityDate, source.IsDeleted, source.DeletedDate, source.AssignedProducts, source.TeamChatMessageCount, source.PrivateChatMessageCount, source.CallCount, source.MeetingCount, source.MeetingsOrganizedCount, source.MeetingsAttendedCount, source.AdHocMeetingsOrganizedCount, source.AdHocMeetingsAttendedCount, source.ScheduledOne-timeMeetingsOrganizedCount, source.ScheduledOne-timeMeetingsAttendedCount, source.ScheduledRecurringMeetingsOrganizedCount, source.ScheduledRecurringMeetingsAttendedCount, source.AudioDuration, source.VideoDuration, source.ScreenShareDuration, source.AudioDurationInSeconds, source.VideoDurationInSeconds, source.ScreenShareDurationInSeconds, source.HasOtherAction, source.UrgentMessages, source.PostMessages, source.TenantDisplayName, source.SharedChannelTenantDisplayNames, source.ReplyMessages, source.IsLicensed, source.ReportPeriod, source.LoadDate); + INSERT (ReportRefreshDate, UserId, UserPrincipalName, LastActivityDate, IsDeleted, DeletedDate, AssignedProducts, TeamChatMessageCount, PrivateChatMessageCount, CallCount, MeetingCount, MeetingsOrganizedCount, MeetingsAttendedCount, AdHocMeetingsOrganizedCount, AdHocMeetingsAttendedCount, [ScheduledOne-timeMeetingsOrganizedCount], [ScheduledOne-timeMeetingsAttendedCount], ScheduledRecurringMeetingsOrganizedCount, ScheduledRecurringMeetingsAttendedCount, AudioDuration, VideoDuration, ScreenShareDuration, AudioDurationInSeconds, VideoDurationInSeconds, ScreenShareDurationInSeconds, HasOtherAction, UrgentMessages, PostMessages, TenantDisplayName, SharedChannelTenantDisplayNames, ReplyMessages, IsLicensed, ReportPeriod, LoadDate) VALUES (source.ReportRefreshDate, source.UserId, source.UserPrincipalName, source.LastActivityDate, source.IsDeleted, source.DeletedDate, source.AssignedProducts, source.TeamChatMessageCount, source.PrivateChatMessageCount, source.CallCount, source.MeetingCount, source.MeetingsOrganizedCount, source.MeetingsAttendedCount, source.AdHocMeetingsOrganizedCount, source.AdHocMeetingsAttendedCount, source.[ScheduledOne-timeMeetingsOrganizedCount], source.[ScheduledOne-timeMeetingsAttendedCount], source.ScheduledRecurringMeetingsOrganizedCount, source.ScheduledRecurringMeetingsAttendedCount, source.AudioDuration, source.VideoDuration, source.ScreenShareDuration, source.AudioDurationInSeconds, source.VideoDurationInSeconds, source.ScreenShareDurationInSeconds, source.HasOtherAction, source.UrgentMessages, source.PostMessages, source.TenantDisplayName, source.SharedChannelTenantDisplayNames, source.ReplyMessages, source.IsLicensed, source.ReportPeriod, source.LoadDate); diff --git a/providers/tests/postgres/dialects/test_postgres.py b/providers/tests/postgres/dialects/test_postgres.py index ab4968a66456b..c7593723325ac 100644 --- a/providers/tests/postgres/dialects/test_postgres.py +++ b/providers/tests/postgres/dialects/test_postgres.py @@ -43,8 +43,9 @@ def get_records(sql, parameters): self.test_db_hook = MagicMock(placeholder="?", inspector=inspector, spec=DbApiHook) self.test_db_hook.get_records.side_effect = get_records - self.test_db_hook._insert_statement_format = "INSERT INTO {} {} VALUES ({})" - self.test_db_hook._escape_column_name_format = '"{}"' + self.test_db_hook.insert_statement_format = "INSERT INTO {} {} VALUES ({})" + self.test_db_hook.escape_word_format = '"{}"' + self.test_db_hook.escape_column_names = False def test_placeholder(self): assert PostgresDialect(self.test_db_hook).placeholder == "?" @@ -69,11 +70,11 @@ def test_get_primary_keys(self): def test_generate_replace_sql(self): values = [ - {"id": "id", "name": "Stallone", "firstname": "Sylvester", "age": "78"}, - {"id": "id", "name": "Statham", "firstname": "Jason", "age": "57"}, - {"id": "id", "name": "Li", "firstname": "Jet", "age": "61"}, - {"id": "id", "name": "Lundgren", "firstname": "Dolph", "age": "66"}, - {"id": "id", "name": "Norris", "firstname": "Chuck", "age": "84"}, + {"id": 1, "name": "Stallone", "firstname": "Sylvester", "age": "78"}, + {"id": 2, "name": "Statham", "firstname": "Jason", "age": "57"}, + {"id": 3, "name": "Li", "firstname": "Jet", "age": "61"}, + {"id": 4, "name": "Lundgren", "firstname": "Dolph", "age": "66"}, + {"id": 5, "name": "Norris", "firstname": "Chuck", "age": "84"}, ] target_fields = ["id", "name", "firstname", "age"] sql = PostgresDialect(self.test_db_hook).generate_replace_sql( @@ -85,3 +86,23 @@ def test_generate_replace_sql(self): INSERT INTO hollywood.actors (id, name, firstname, age) VALUES (?,?,?,?,?) ON CONFLICT (id) DO UPDATE SET name = excluded.name, firstname = excluded.firstname, age = excluded.age """.strip() ) + + def test_generate_replace_sql_when_escape_column_names_is_enabled(self): + values = [ + {"id": 1, "name": "Stallone", "firstname": "Sylvester", "age": "78"}, + {"id": 2, "name": "Statham", "firstname": "Jason", "age": "57"}, + {"id": 3, "name": "Li", "firstname": "Jet", "age": "61"}, + {"id": 4, "name": "Lundgren", "firstname": "Dolph", "age": "66"}, + {"id": 5, "name": "Norris", "firstname": "Chuck", "age": "84"}, + ] + target_fields = ["id", "name", "firstname", "age"] + self.test_db_hook.escape_column_names = True + sql = PostgresDialect(self.test_db_hook).generate_replace_sql( + "hollywood.actors", values, target_fields + ) + assert ( + sql + == """ + INSERT INTO hollywood.actors ("id", "name", "firstname", "age") VALUES (?,?,?,?,?) ON CONFLICT ("id") DO UPDATE SET "name" = excluded."name", "firstname" = excluded."firstname", "age" = excluded."age" + """.strip() + ) diff --git a/providers/tests/teradata/hooks/test_teradata.py b/providers/tests/teradata/hooks/test_teradata.py index 10754555b1a4d..f9a38e0d607f0 100644 --- a/providers/tests/teradata/hooks/test_teradata.py +++ b/providers/tests/teradata/hooks/test_teradata.py @@ -41,6 +41,7 @@ def setup_method(self): self.cur = mock.MagicMock(rowcount=0) self.conn = mock.MagicMock() self.conn.cursor.return_value = self.cur + self.conn.extra_dejson = {} conn = self.conn class UnitTestTeradataHook(TeradataHook):