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..e94bdffa52a3 100644 --- a/packages/sqlalchemy-spanner/google/cloud/sqlalchemy_spanner/sqlalchemy_spanner.py +++ b/packages/sqlalchemy-spanner/google/cloud/sqlalchemy_spanner/sqlalchemy_spanner.py @@ -25,7 +25,7 @@ ) from google.api_core.client_options import ClientOptions from google.auth.credentials import AnonymousCredentials -from google.cloud.spanner_v1 import Client, TransactionOptions +from google.cloud.spanner_v1 import Client, TransactionOptions, param_types from google.cloud.spanner_v1.data_types import JsonObject from sqlalchemy import ForeignKeyConstraint, PickleType, TypeDecorator, types from sqlalchemy.engine.base import Engine @@ -967,22 +967,36 @@ def _get_table_filter_query( ): """ Generates WHERE query for tables or views for which - information is reflected. + information is reflected. The names are bound as the + ``@filter_names`` query parameter, see ``_get_reflection_params``. """ table_filter_query = "" if filter_names is not None: - for table_name in filter_names: - query = f"{info_schema_table}.table_name = '{table_name}'" - if table_filter_query != "": - table_filter_query = table_filter_query + " OR " + query - else: - table_filter_query = query - table_filter_query = "(" + table_filter_query + ") " + table_filter_query = ( + f"({info_schema_table}.table_name IN UNNEST(@filter_names)) " + ) if append_query: table_filter_query = table_filter_query + " AND " return table_filter_query + def _get_reflection_params(self, schema, filter_names=None, **names): + """ + Generates the query parameters and parameter types for the + INFORMATION_SCHEMA reflection queries. The schema, the optional + list of table names and any additional names are bound as query + parameters instead of being interpolated into the SQL text. + """ + params = {"schema": schema or ""} + types = {"schema": param_types.STRING} + for name, value in names.items(): + params[name] = value + types[name] = param_types.STRING + if filter_names is not None: + params["filter_names"] = list(filter_names) + types["filter_names"] = param_types.Array(param_types.STRING) + return params, types + def create_connect_args(self, url): """Parse connection args from the given URL. @@ -1045,12 +1059,13 @@ def get_view_names(self, connection, schema=None, **kw): sql = """ SELECT table_name FROM information_schema.views - WHERE TABLE_SCHEMA='{}' - """.format(schema or "") + WHERE TABLE_SCHEMA=@schema + """ + params, types = self._get_reflection_params(schema) all_views = [] with connection.connection.database.snapshot() as snap: - rows = list(snap.execute_sql(sql)) + rows = list(snap.execute_sql(sql, params=params, param_types=types)) for view in rows: all_views.append(view[0]) @@ -1074,11 +1089,12 @@ def get_sequence_names(self, connection, schema=None, **kw): sql = """ SELECT name FROM information_schema.sequences - WHERE SCHEMA='{}' - """.format(schema or "") + WHERE SCHEMA=@schema + """ + params, types = self._get_reflection_params(schema) all_sequences = [] with connection.connection.database.snapshot() as snap: - rows = list(snap.execute_sql(sql)) + rows = list(snap.execute_sql(sql, params=params, param_types=types)) for seq in rows: all_sequences.append(seq[0]) @@ -1103,11 +1119,12 @@ def get_view_definition(self, connection, view_name, schema=None, **kw): sql = """ 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) + WHERE TABLE_SCHEMA=@schema AND TABLE_NAME=@view_name + """ + params, types = self._get_reflection_params(schema, view_name=view_name) with connection.connection.database.snapshot() as snap: - rows = list(snap.execute_sql(sql)) + rows = list(snap.execute_sql(sql, params=params, param_types=types)) if rows == []: raise NoSuchTableError(f"{schema if schema else ''}.{view_name}") result = rows[0][0] @@ -1143,10 +1160,9 @@ def get_multi_columns( The schema is ``None`` if no schema is provided. """ 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_filter_query = " col.table_schema = @schema AND " table_type_query = self._get_table_type_query(kind, True) + params, types = self._get_reflection_params(schema, filter_names) sql = """ SELECT col.table_schema, col.table_name, col.column_name, @@ -1171,7 +1187,7 @@ def get_multi_columns( schema_filter_query=schema_filter_query, ) with connection.connection.database.snapshot() as snap: - columns = list(snap.execute_sql(sql)) + columns = list(snap.execute_sql(sql, params=params, param_types=types)) result_dict = {} for col in columns: @@ -1272,10 +1288,9 @@ def get_multi_indexes( The schema is ``None`` if no schema is provided. """ 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_filter_query = " i.table_schema = @schema AND " table_type_query = self._get_table_type_query(kind, True) + params, types = self._get_reflection_params(schema, filter_names) sql = """ SELECT @@ -1335,7 +1350,7 @@ def get_multi_indexes( ) with connection.connection.database.snapshot() as snap: - rows = list(snap.execute_sql(sql)) + rows = list(snap.execute_sql(sql, params=params, param_types=types)) result_dict = {} for row in rows: @@ -1413,10 +1428,9 @@ def get_multi_pk_constraint( The schema is ``None`` if no schema is provided. """ 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_filter_query = " tc.table_schema = @schema AND " table_type_query = self._get_table_type_query(kind, True) + params, types = self._get_reflection_params(schema, filter_names) sql = """ SELECT tc.table_schema, tc.table_name, kcu.column_name @@ -1438,7 +1452,7 @@ def get_multi_pk_constraint( ) with connection.connection.database.snapshot() as snap: - rows = list(snap.execute_sql(sql)) + rows = list(snap.execute_sql(sql, params=params, param_types=types)) result_dict = {} for row in rows: @@ -1524,10 +1538,9 @@ def get_multi_foreign_keys( The schema is ``None`` if no schema is provided. """ 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_filter_query = " tc.table_schema = @schema AND " table_type_query = self._get_table_type_query(kind, True) + params, types = self._get_reflection_params(schema, filter_names) sql = """ SELECT @@ -1580,7 +1593,7 @@ def get_multi_foreign_keys( ) with connection.connection.database.snapshot() as snap: - rows = list(snap.execute_sql(sql)) + rows = list(snap.execute_sql(sql, params=params, param_types=types)) result_dict = {} for row in rows: @@ -1640,12 +1653,13 @@ def get_table_names(self, connection, schema=None, **kw): sql = """ SELECT table_name FROM information_schema.tables -WHERE table_type = 'BASE TABLE' AND table_schema = '{schema}' -""".format(schema=schema or "") +WHERE table_type = 'BASE TABLE' AND table_schema = @schema +""" + params, types = self._get_reflection_params(schema) table_names = [] with connection.connection.database.snapshot() as snap: - rows = snap.execute_sql(sql) + rows = snap.execute_sql(sql, params=params, param_types=types) for row in rows: table_names.append(row[0]) @@ -1673,15 +1687,16 @@ def get_unique_constraints(self, connection, table_name, schema=None, **kw): JOIN INFORMATION_SCHEMA.CONSTRAINT_COLUMN_USAGE AS ccu USING (TABLE_CATALOG, TABLE_SCHEMA, CONSTRAINT_NAME) WHERE - tc.TABLE_NAME="{table_name}" - AND tc.TABLE_SCHEMA="{table_schema}" + tc.TABLE_NAME=@table_name + AND tc.TABLE_SCHEMA=@schema AND tc.CONSTRAINT_TYPE = "UNIQUE" AND tc.CONSTRAINT_NAME IS NOT NULL -""".format(table_schema=schema or "", table_name=table_name) +""" + params, types = self._get_reflection_params(schema, table_name=table_name) cols = [] with connection.connection.database.snapshot() as snap: - rows = snap.execute_sql(sql) + rows = snap.execute_sql(sql, params=params, param_types=types) for row in rows: cols.append({"name": row[0], "column_names": [row[1]]}) @@ -1703,14 +1718,17 @@ def has_table(self, connection, table_name, schema=None, **kw): Returns: bool: True, if the given table exists, False otherwise. """ + params, types = self._get_reflection_params(schema, table_name=table_name) with connection.connection.database.snapshot() as snap: rows = snap.execute_sql( """ SELECT true FROM INFORMATION_SCHEMA.TABLES -WHERE TABLE_SCHEMA="{table_schema}" AND TABLE_NAME="{table_name}" +WHERE TABLE_SCHEMA=@schema AND TABLE_NAME=@table_name LIMIT 1 -""".format(table_schema=schema or "", table_name=table_name) +""", + params=params, + param_types=types, ) for _ in rows: @@ -1727,15 +1745,18 @@ def has_sequence(self, connection, sequence_name, schema=None, **kw): the database, False otherwise. """ + params, types = self._get_reflection_params(schema, sequence_name=sequence_name) with connection.connection.database.snapshot() as snap: rows = snap.execute_sql( """ SELECT true FROM INFORMATION_SCHEMA.SEQUENCES - WHERE NAME="{sequence_name}" - AND SCHEMA="{schema}" + WHERE NAME=@sequence_name + AND SCHEMA=@schema LIMIT 1 - """.format(sequence_name=sequence_name, schema=schema or "") + """, + params=params, + param_types=types, ) for _ in rows: diff --git a/packages/sqlalchemy-spanner/tests/mockserver_tests/test_auto_increment.py b/packages/sqlalchemy-spanner/tests/mockserver_tests/test_auto_increment.py index 72adceba7fbd..1ade4f7b8e38 100644 --- a/packages/sqlalchemy-spanner/tests/mockserver_tests/test_auto_increment.py +++ b/packages/sqlalchemy-spanner/tests/mockserver_tests/test_auto_increment.py @@ -38,7 +38,7 @@ def test_create_table(self): add_result( """SELECT true FROM INFORMATION_SCHEMA.TABLES -WHERE TABLE_SCHEMA="" AND TABLE_NAME="singers" +WHERE TABLE_SCHEMA=@schema AND TABLE_NAME=@table_name LIMIT 1 """, ResultSet(), @@ -64,7 +64,7 @@ def test_create_auto_increment_table(self): add_result( """SELECT true FROM INFORMATION_SCHEMA.TABLES -WHERE TABLE_SCHEMA="" AND TABLE_NAME="singers" +WHERE TABLE_SCHEMA=@schema AND TABLE_NAME=@table_name LIMIT 1 """, ResultSet(), @@ -90,7 +90,7 @@ def test_create_table_with_specific_sequence_kind(self): add_result( """SELECT true FROM INFORMATION_SCHEMA.TABLES -WHERE TABLE_SCHEMA="" AND TABLE_NAME="singers" +WHERE TABLE_SCHEMA=@schema AND TABLE_NAME=@table_name LIMIT 1 """, ResultSet(), diff --git a/packages/sqlalchemy-spanner/tests/mockserver_tests/test_basics.py b/packages/sqlalchemy-spanner/tests/mockserver_tests/test_basics.py index bed297567388..ab5f3cabcbc2 100644 --- a/packages/sqlalchemy-spanner/tests/mockserver_tests/test_basics.py +++ b/packages/sqlalchemy-spanner/tests/mockserver_tests/test_basics.py @@ -100,7 +100,7 @@ def test_create_table(self): add_result( """SELECT true FROM INFORMATION_SCHEMA.TABLES -WHERE TABLE_SCHEMA="" AND TABLE_NAME="users" +WHERE TABLE_SCHEMA=@schema AND TABLE_NAME=@table_name LIMIT 1 """, ResultSet(), @@ -134,7 +134,7 @@ def test_create_table_in_schema(self): add_result( """SELECT true FROM INFORMATION_SCHEMA.TABLES -WHERE TABLE_SCHEMA="schema" AND TABLE_NAME="users" +WHERE TABLE_SCHEMA=@schema AND TABLE_NAME=@table_name LIMIT 1 """, ResultSet(), @@ -169,15 +169,14 @@ def test_create_table_in_schema(self): ) def test_create_multiple_tables(self): - for i in range(2): - add_result( - f"""SELECT true + add_result( + """SELECT true FROM INFORMATION_SCHEMA.TABLES -WHERE TABLE_SCHEMA="" AND TABLE_NAME="table{i}" +WHERE TABLE_SCHEMA=@schema AND TABLE_NAME=@table_name LIMIT 1 """, - ResultSet(), - ) + ResultSet(), + ) engine = self.create_engine() metadata = MetaData() for i in range(2): diff --git a/packages/sqlalchemy-spanner/tests/mockserver_tests/test_bit_reversed_sequence.py b/packages/sqlalchemy-spanner/tests/mockserver_tests/test_bit_reversed_sequence.py index 97553689f6f1..b7102724157f 100644 --- a/packages/sqlalchemy-spanner/tests/mockserver_tests/test_bit_reversed_sequence.py +++ b/packages/sqlalchemy-spanner/tests/mockserver_tests/test_bit_reversed_sequence.py @@ -37,7 +37,7 @@ def test_create_table(self): add_result( """SELECT true FROM INFORMATION_SCHEMA.TABLES -WHERE TABLE_SCHEMA="" AND TABLE_NAME="singers" +WHERE TABLE_SCHEMA=@schema AND TABLE_NAME=@table_name LIMIT 1 """, ResultSet(), @@ -45,8 +45,8 @@ def test_create_table(self): add_result( """SELECT true FROM INFORMATION_SCHEMA.SEQUENCES - WHERE NAME="singer_id" - AND SCHEMA="" + WHERE NAME=@sequence_name + AND SCHEMA=@schema LIMIT 1""", ResultSet(), ) diff --git a/packages/sqlalchemy-spanner/tests/mockserver_tests/test_commit_timestamp.py b/packages/sqlalchemy-spanner/tests/mockserver_tests/test_commit_timestamp.py index ee77498ebe89..12563950cf32 100644 --- a/packages/sqlalchemy-spanner/tests/mockserver_tests/test_commit_timestamp.py +++ b/packages/sqlalchemy-spanner/tests/mockserver_tests/test_commit_timestamp.py @@ -29,7 +29,7 @@ def test_create_table(self): add_result( """SELECT true FROM INFORMATION_SCHEMA.TABLES -WHERE TABLE_SCHEMA="" AND TABLE_NAME="singers" +WHERE TABLE_SCHEMA=@schema AND TABLE_NAME=@table_name LIMIT 1 """, ResultSet(), @@ -37,8 +37,8 @@ def test_create_table(self): add_result( """SELECT true FROM INFORMATION_SCHEMA.SEQUENCES - WHERE NAME="singer_id" - AND SCHEMA="" + WHERE NAME=@sequence_name + AND SCHEMA=@schema LIMIT 1""", ResultSet(), ) diff --git a/packages/sqlalchemy-spanner/tests/mockserver_tests/test_default.py b/packages/sqlalchemy-spanner/tests/mockserver_tests/test_default.py index da2258f34a48..889ea84d68f2 100644 --- a/packages/sqlalchemy-spanner/tests/mockserver_tests/test_default.py +++ b/packages/sqlalchemy-spanner/tests/mockserver_tests/test_default.py @@ -27,7 +27,7 @@ def test_create_table_with_default(self): add_result( """SELECT true FROM INFORMATION_SCHEMA.TABLES -WHERE TABLE_SCHEMA="" AND TABLE_NAME="singers" +WHERE TABLE_SCHEMA=@schema AND TABLE_NAME=@table_name LIMIT 1 """, ResultSet(), diff --git a/packages/sqlalchemy-spanner/tests/mockserver_tests/test_interleaved_index.py b/packages/sqlalchemy-spanner/tests/mockserver_tests/test_interleaved_index.py index 14f362036734..fc4cb676b2db 100644 --- a/packages/sqlalchemy-spanner/tests/mockserver_tests/test_interleaved_index.py +++ b/packages/sqlalchemy-spanner/tests/mockserver_tests/test_interleaved_index.py @@ -35,7 +35,7 @@ def test_create_table(self): add_result( """SELECT true FROM INFORMATION_SCHEMA.TABLES -WHERE TABLE_SCHEMA="" AND TABLE_NAME="singers" +WHERE TABLE_SCHEMA=@schema AND TABLE_NAME=@table_name LIMIT 1 """, ResultSet(), @@ -43,7 +43,7 @@ def test_create_table(self): add_result( """SELECT true FROM INFORMATION_SCHEMA.TABLES -WHERE TABLE_SCHEMA="" AND TABLE_NAME="albums" +WHERE TABLE_SCHEMA=@schema AND TABLE_NAME=@table_name LIMIT 1 """, ResultSet(), @@ -51,7 +51,7 @@ def test_create_table(self): add_result( """SELECT true FROM INFORMATION_SCHEMA.TABLES -WHERE TABLE_SCHEMA="" AND TABLE_NAME="tracks" +WHERE TABLE_SCHEMA=@schema AND TABLE_NAME=@table_name LIMIT 1 """, ResultSet(), diff --git a/packages/sqlalchemy-spanner/tests/mockserver_tests/test_json.py b/packages/sqlalchemy-spanner/tests/mockserver_tests/test_json.py index 2cec82a2ebc3..f44839dfa800 100644 --- a/packages/sqlalchemy-spanner/tests/mockserver_tests/test_json.py +++ b/packages/sqlalchemy-spanner/tests/mockserver_tests/test_json.py @@ -41,7 +41,7 @@ def test_create_table(self): add_result( """SELECT true FROM INFORMATION_SCHEMA.TABLES -WHERE TABLE_SCHEMA="" AND TABLE_NAME="venues" +WHERE TABLE_SCHEMA=@schema AND TABLE_NAME=@table_name LIMIT 1 """, ResultSet(), diff --git a/packages/sqlalchemy-spanner/tests/mockserver_tests/test_not_enforced_fk.py b/packages/sqlalchemy-spanner/tests/mockserver_tests/test_not_enforced_fk.py index bad9d5b7f669..bce3ae432a37 100644 --- a/packages/sqlalchemy-spanner/tests/mockserver_tests/test_not_enforced_fk.py +++ b/packages/sqlalchemy-spanner/tests/mockserver_tests/test_not_enforced_fk.py @@ -35,7 +35,7 @@ def test_create_table(self): add_result( """SELECT true FROM INFORMATION_SCHEMA.TABLES -WHERE TABLE_SCHEMA="" AND TABLE_NAME="singers" +WHERE TABLE_SCHEMA=@schema AND TABLE_NAME=@table_name LIMIT 1 """, ResultSet(), @@ -43,7 +43,7 @@ def test_create_table(self): add_result( """SELECT true FROM INFORMATION_SCHEMA.TABLES -WHERE TABLE_SCHEMA="" AND TABLE_NAME="albums" +WHERE TABLE_SCHEMA=@schema AND TABLE_NAME=@table_name LIMIT 1 """, ResultSet(), diff --git a/packages/sqlalchemy-spanner/tests/mockserver_tests/test_null_filtered_index.py b/packages/sqlalchemy-spanner/tests/mockserver_tests/test_null_filtered_index.py index e5d8de3c43f8..56253188d2c4 100644 --- a/packages/sqlalchemy-spanner/tests/mockserver_tests/test_null_filtered_index.py +++ b/packages/sqlalchemy-spanner/tests/mockserver_tests/test_null_filtered_index.py @@ -35,7 +35,7 @@ def test_create_table(self): add_result( """SELECT true FROM INFORMATION_SCHEMA.TABLES -WHERE TABLE_SCHEMA="" AND TABLE_NAME="singers" +WHERE TABLE_SCHEMA=@schema AND TABLE_NAME=@table_name LIMIT 1 """, ResultSet(), @@ -43,7 +43,7 @@ def test_create_table(self): add_result( """SELECT true FROM INFORMATION_SCHEMA.TABLES -WHERE TABLE_SCHEMA="" AND TABLE_NAME="albums" +WHERE TABLE_SCHEMA=@schema AND TABLE_NAME=@table_name LIMIT 1 """, ResultSet(), diff --git a/packages/sqlalchemy-spanner/tests/mockserver_tests/test_pickle_type.py b/packages/sqlalchemy-spanner/tests/mockserver_tests/test_pickle_type.py index 3dbef8e3621f..421e36fb6090 100644 --- a/packages/sqlalchemy-spanner/tests/mockserver_tests/test_pickle_type.py +++ b/packages/sqlalchemy-spanner/tests/mockserver_tests/test_pickle_type.py @@ -39,7 +39,7 @@ def test_create_table(self): add_result( """SELECT true FROM INFORMATION_SCHEMA.TABLES -WHERE TABLE_SCHEMA="" AND TABLE_NAME="user_preferences" +WHERE TABLE_SCHEMA=@schema AND TABLE_NAME=@table_name LIMIT 1 """, ResultSet(), diff --git a/packages/sqlalchemy-spanner/tests/mockserver_tests/test_quickstart.py b/packages/sqlalchemy-spanner/tests/mockserver_tests/test_quickstart.py index cb8497c81206..cedbf3f49daa 100644 --- a/packages/sqlalchemy-spanner/tests/mockserver_tests/test_quickstart.py +++ b/packages/sqlalchemy-spanner/tests/mockserver_tests/test_quickstart.py @@ -33,14 +33,14 @@ def test_create_tables(self): add_result( """SELECT true FROM INFORMATION_SCHEMA.TABLES -WHERE TABLE_SCHEMA="" AND TABLE_NAME="user_account" +WHERE TABLE_SCHEMA=@schema AND TABLE_NAME=@table_name LIMIT 1""", ResultSet(), ) add_result( """SELECT true FROM INFORMATION_SCHEMA.TABLES -WHERE TABLE_SCHEMA="" AND TABLE_NAME="address" +WHERE TABLE_SCHEMA=@schema AND TABLE_NAME=@table_name LIMIT 1""", ResultSet(), ) diff --git a/packages/sqlalchemy-spanner/tests/unit/test_dialect.py b/packages/sqlalchemy-spanner/tests/unit/test_dialect.py index f44599bb4fcb..5a5726e659be 100644 --- a/packages/sqlalchemy-spanner/tests/unit/test_dialect.py +++ b/packages/sqlalchemy-spanner/tests/unit/test_dialect.py @@ -14,6 +14,7 @@ from unittest.mock import MagicMock +from google.cloud.spanner_v1 import param_types from sqlalchemy.testing import eq_ from sqlalchemy.testing.plugin.plugin_base import fixtures @@ -100,3 +101,98 @@ 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_binds_names_as_query_parameters(self): + """Table and schema names are bound as query parameters instead of + being interpolated into the INFORMATION_SCHEMA query.""" + dialect = SpannerDialect() + connection, mock_snapshot = self._mock_connection() + + dialect.get_columns(connection, table_name="t' OR '1'='1", schema="s") + + sql = mock_snapshot.execute_sql.call_args[0][0] + kwargs = mock_snapshot.execute_sql.call_args[1] + assert "col.table_name IN UNNEST(@filter_names)" in sql + assert "col.table_schema = @schema AND" in sql + assert "'1'='1'" not in sql + eq_(kwargs["params"], {"schema": "s", "filter_names": ["t' OR '1'='1"]}) + eq_( + kwargs["param_types"], + { + "schema": param_types.STRING, + "filter_names": param_types.Array(param_types.STRING), + }, + ) + + def test_get_multi_columns_without_filter_names(self): + """Without filter names no table filter is added and only the schema + is bound as a query parameter.""" + dialect = SpannerDialect() + connection, mock_snapshot = self._mock_connection() + + dialect.get_multi_columns(connection) + + sql = mock_snapshot.execute_sql.call_args[0][0] + kwargs = mock_snapshot.execute_sql.call_args[1] + assert "@filter_names" not in sql + assert "col.table_schema = @schema AND" in sql + eq_(kwargs["params"], {"schema": ""}) + eq_(kwargs["param_types"], {"schema": param_types.STRING}) + + def test_has_table_binds_names_as_query_parameters(self): + dialect = SpannerDialect() + connection, mock_snapshot = self._mock_connection() + + eq_(dialect.has_table(connection, table_name='a" OR "1"="1'), False) + + sql = mock_snapshot.execute_sql.call_args[0][0] + kwargs = mock_snapshot.execute_sql.call_args[1] + assert "WHERE TABLE_SCHEMA=@schema AND TABLE_NAME=@table_name" in sql + assert '"1"="1"' not in sql + eq_(kwargs["params"], {"schema": "", "table_name": 'a" OR "1"="1'}) + eq_( + kwargs["param_types"], + {"schema": param_types.STRING, "table_name": param_types.STRING}, + ) + + def test_get_view_definition_binds_names_as_query_parameters(self): + dialect = SpannerDialect() + connection, mock_snapshot = self._mock_connection(rows=[["SELECT 1"]]) + + definition = dialect.get_view_definition(connection, view_name="v", schema="s") + + eq_(definition, "SELECT 1") + sql = mock_snapshot.execute_sql.call_args[0][0] + kwargs = mock_snapshot.execute_sql.call_args[1] + assert "WHERE TABLE_SCHEMA=@schema AND TABLE_NAME=@view_name" in sql + eq_(kwargs["params"], {"schema": "s", "view_name": "v"}) + eq_( + kwargs["param_types"], + {"schema": param_types.STRING, "view_name": param_types.STRING}, + ) + + def test_has_sequence_binds_names_as_query_parameters(self): + dialect = SpannerDialect() + connection, mock_snapshot = self._mock_connection(rows=[[True]]) + + eq_(dialect.has_sequence(connection, sequence_name="seq"), True) + + sql = mock_snapshot.execute_sql.call_args[0][0] + kwargs = mock_snapshot.execute_sql.call_args[1] + assert "WHERE NAME=@sequence_name" in sql + assert "AND SCHEMA=@schema" in sql + eq_(kwargs["params"], {"schema": "", "sequence_name": "seq"}) + eq_( + kwargs["param_types"], + {"schema": param_types.STRING, "sequence_name": param_types.STRING}, + )