From 0411049c63203990c8a822022ca4f05b32381386 Mon Sep 17 00:00:00 2001 From: bibi samina Date: Mon, 24 Aug 2026 14:33:54 +0530 Subject: [PATCH] fix(sqlalchemy-spanner): escape names in reflection queries --- .../sqlalchemy_spanner/sqlalchemy_spanner.py | 54 ++++++++++++++---- .../tests/unit/test_dialect.py | 56 +++++++++++++++++++ 2 files changed, 100 insertions(+), 10 deletions(-) diff --git a/packages/sqlalchemy-spanner/google/cloud/sqlalchemy_spanner/sqlalchemy_spanner.py b/packages/sqlalchemy-spanner/google/cloud/sqlalchemy_spanner/sqlalchemy_spanner.py index ee8e72eb5665..43ad0eccfe97 100644 --- a/packages/sqlalchemy-spanner/google/cloud/sqlalchemy_spanner/sqlalchemy_spanner.py +++ b/packages/sqlalchemy-spanner/google/cloud/sqlalchemy_spanner/sqlalchemy_spanner.py @@ -76,6 +76,25 @@ def reset_connection(dbapi_conn, connection_record, reset_state=None): OPERATORS[json_getitem_op] = operator_lookup["json_getitem_op"] +def _escape_sql_string_literal(value): + """Escape a value for safe inclusion in a GoogleSQL string literal. + + The reflection queries below build ``INFORMATION_SCHEMA`` predicates by + interpolating table, schema, view and sequence names into quoted string + literals. A name containing a quote (for example, one enumerated from a + shared or foreign database and fed back in during reflection) would + otherwise close the literal so the remainder is parsed as SQL. Escaping the + backslash, both quote characters and newlines keeps the name contained. + """ + return ( + value.replace("\\", "\\\\") + .replace("'", "\\'") + .replace('"', '\\"') + .replace("\n", "\\n") + .replace("\r", "\\r") + ) + + # PickleType that can be used with Spanner. # Binary values are automatically encoded/decoded to/from base64. # Usage: @@ -972,7 +991,10 @@ def _get_table_filter_query( table_filter_query = "" if filter_names is not None: for table_name in filter_names: - query = f"{info_schema_table}.table_name = '{table_name}'" + query = ( + f"{info_schema_table}.table_name = " + f"'{_escape_sql_string_literal(table_name)}'" + ) if table_filter_query != "": table_filter_query = table_filter_query + " OR " + query else: @@ -1104,7 +1126,10 @@ def get_view_definition(self, connection, view_name, schema=None, **kw): SELECT view_definition FROM information_schema.views WHERE TABLE_SCHEMA='{schema_name}' AND TABLE_NAME='{view_name}' - """.format(schema_name=schema or "", view_name=view_name) + """.format( + schema_name=_escape_sql_string_literal(schema or ""), + view_name=_escape_sql_string_literal(view_name), + ) with connection.connection.database.snapshot() as snap: rows = list(snap.execute_sql(sql)) @@ -1144,7 +1169,7 @@ def get_multi_columns( """ table_filter_query = self._get_table_filter_query(filter_names, "col", True) schema_filter_query = " col.table_schema = '{schema}' AND ".format( - schema=schema or "" + schema=_escape_sql_string_literal(schema or "") ) table_type_query = self._get_table_type_query(kind, True) @@ -1273,7 +1298,7 @@ def get_multi_indexes( """ table_filter_query = self._get_table_filter_query(filter_names, "i", True) schema_filter_query = " i.table_schema = '{schema}' AND ".format( - schema=schema or "" + schema=_escape_sql_string_literal(schema or "") ) table_type_query = self._get_table_type_query(kind, True) @@ -1414,7 +1439,7 @@ def get_multi_pk_constraint( """ table_filter_query = self._get_table_filter_query(filter_names, "tc", True) schema_filter_query = " tc.table_schema = '{schema}' AND ".format( - schema=schema or "" + schema=_escape_sql_string_literal(schema or "") ) table_type_query = self._get_table_type_query(kind, True) @@ -1525,7 +1550,7 @@ def get_multi_foreign_keys( """ table_filter_query = self._get_table_filter_query(filter_names, "tc", True) schema_filter_query = " tc.table_schema = '{schema}' AND".format( - schema=schema or "" + schema=_escape_sql_string_literal(schema or "") ) table_type_query = self._get_table_type_query(kind, True) @@ -1641,7 +1666,7 @@ def get_table_names(self, connection, schema=None, **kw): SELECT table_name FROM information_schema.tables WHERE table_type = 'BASE TABLE' AND table_schema = '{schema}' -""".format(schema=schema or "") +""".format(schema=_escape_sql_string_literal(schema or "")) table_names = [] with connection.connection.database.snapshot() as snap: @@ -1677,7 +1702,10 @@ def get_unique_constraints(self, connection, table_name, schema=None, **kw): AND tc.TABLE_SCHEMA="{table_schema}" AND tc.CONSTRAINT_TYPE = "UNIQUE" AND tc.CONSTRAINT_NAME IS NOT NULL -""".format(table_schema=schema or "", table_name=table_name) +""".format( + table_schema=_escape_sql_string_literal(schema or ""), + table_name=_escape_sql_string_literal(table_name), + ) cols = [] with connection.connection.database.snapshot() as snap: @@ -1710,7 +1738,10 @@ def has_table(self, connection, table_name, schema=None, **kw): FROM INFORMATION_SCHEMA.TABLES WHERE TABLE_SCHEMA="{table_schema}" AND TABLE_NAME="{table_name}" LIMIT 1 -""".format(table_schema=schema or "", table_name=table_name) +""".format( + table_schema=_escape_sql_string_literal(schema or ""), + table_name=_escape_sql_string_literal(table_name), + ) ) for _ in rows: @@ -1735,7 +1766,10 @@ def has_sequence(self, connection, sequence_name, schema=None, **kw): WHERE NAME="{sequence_name}" AND SCHEMA="{schema}" LIMIT 1 - """.format(sequence_name=sequence_name, schema=schema or "") + """.format( + sequence_name=_escape_sql_string_literal(sequence_name), + schema=_escape_sql_string_literal(schema or ""), + ) ) for _ in rows: diff --git a/packages/sqlalchemy-spanner/tests/unit/test_dialect.py b/packages/sqlalchemy-spanner/tests/unit/test_dialect.py index 86e0907137f1..44b2fa79e75f 100644 --- a/packages/sqlalchemy-spanner/tests/unit/test_dialect.py +++ b/packages/sqlalchemy-spanner/tests/unit/test_dialect.py @@ -98,3 +98,59 @@ def test_max_size_exported(self): eq_(SpannerDialect.max_size, MAX_SIZE) eq_(int_from_size("MAX"), 2621440) eq_(int_from_size("100"), 100) + + @staticmethod + def _mock_connection(rows=None): + connection = MagicMock() + mock_snapshot = MagicMock() + mock_snapshot.execute_sql.return_value = rows if rows is not None else [] + connection.connection.database.snapshot.return_value.__enter__.return_value = ( + mock_snapshot + ) + return connection, mock_snapshot + + def test_get_columns_escapes_quote_in_table_name(self): + """A single quote in a reflected table name must not break out of the + INFORMATION_SCHEMA string literal in get_columns.""" + dialect = SpannerDialect() + connection, mock_snapshot = self._mock_connection() + + dialect.get_columns(connection, table_name="t' OR '1'='1") + + executed_sql = mock_snapshot.execute_sql.call_args[0][0] + assert "col.table_name = 't\\' OR \\'1\\'=\\'1'" in executed_sql + assert "col.table_name = 't' OR '1'='1'" not in executed_sql + + def test_has_table_escapes_quote_in_table_name(self): + """A double quote in a reflected table name must not break out of the + INFORMATION_SCHEMA string literal in has_table.""" + dialect = SpannerDialect() + connection, mock_snapshot = self._mock_connection() + + dialect.has_table(connection, table_name='a" OR "1"="1') + + executed_sql = mock_snapshot.execute_sql.call_args[0][0] + assert 'TABLE_NAME="a\\" OR \\"1\\"=\\"1"' in executed_sql + assert 'TABLE_NAME="a" OR "1"="1"' not in executed_sql + + def test_get_view_definition_escapes_quote(self): + """A quote in a reflected view name must not break out of the literal.""" + dialect = SpannerDialect() + connection, mock_snapshot = self._mock_connection(rows=[["def"]]) + + dialect.get_view_definition(connection, view_name="v' OR '1'='1") + + executed_sql = mock_snapshot.execute_sql.call_args[0][0] + assert "TABLE_NAME='v\\' OR \\'1\\'=\\'1'" in executed_sql + + def test_escape_sql_string_literal(self): + """The helper escapes backslashes, both quote styles and newlines.""" + from google.cloud.sqlalchemy_spanner.sqlalchemy_spanner import ( + _escape_sql_string_literal, + ) + + eq_(_escape_sql_string_literal("a'b"), "a\\'b") + eq_(_escape_sql_string_literal('a"b'), 'a\\"b') + eq_(_escape_sql_string_literal("a\\b"), "a\\\\b") + eq_(_escape_sql_string_literal("a\nb"), "a\\nb") + eq_(_escape_sql_string_literal("plain"), "plain")