diff --git a/packages/django-google-spanner/.coveragerc b/packages/django-google-spanner/.coveragerc index fd5adc4f59da..13735bf8f157 100644 --- a/packages/django-google-spanner/.coveragerc +++ b/packages/django-google-spanner/.coveragerc @@ -10,14 +10,13 @@ branch = True source = django_spanner - [paths] source = django_spanner */site-packages/django_spanner [report] -fail_under = 68 +fail_under = 100 show_missing = True exclude_lines = # Re-enable the standard pragma @@ -26,6 +25,7 @@ exclude_lines = def __repr__ # Ignore abstract methods raise NotImplementedError + if __name__ == .__main__.: omit = tests/* */tests/* \ No newline at end of file diff --git a/packages/django-google-spanner/tests/mockserver_tests/test_basics.py b/packages/django-google-spanner/tests/mockserver_tests/test_basics.py index 8d53e081a745..203bba887456 100644 --- a/packages/django-google-spanner/tests/mockserver_tests/test_basics.py +++ b/packages/django-google-spanner/tests/mockserver_tests/test_basics.py @@ -132,4 +132,4 @@ class LocalSinger(models.Model): finally: for db, config in DATABASES.items(): if config["ENGINE"] == "django_spanner": - config.pop("DISABLE_RANDOM_ID_GENERATION", None) + config.pop("RANDOM_ID_GENERATION_ENABLED", None) diff --git a/packages/django-google-spanner/tests/unit/django_spanner/test__opentelemetry_tracing.py b/packages/django-google-spanner/tests/unit/django_spanner/test__opentelemetry_tracing.py index e39689e6c276..bd4daa803762 100644 --- a/packages/django-google-spanner/tests/unit/django_spanner/test__opentelemetry_tracing.py +++ b/packages/django-google-spanner/tests/unit/django_spanner/test__opentelemetry_tracing.py @@ -100,6 +100,12 @@ def test_trace_call(self): self.assertEqual(span.name, "CloudSpannerDjango.Test") self.assertEqual(span.status.status_code, StatusCode.OK) + def test_trace_call_no_extra_attributes(self): + with _opentelemetry_tracing.trace_call( + "CloudSpannerDjango.TestNoExtra", _make_connection() + ) as span: + self.assertIsNotNone(span) + def test_trace_error(self): extra_attributes = {"db.instance": "database_name"} diff --git a/packages/django-google-spanner/tests/unit/django_spanner/test_base.py b/packages/django-google-spanner/tests/unit/django_spanner/test_base.py index f6595688646b..a4309db0e027 100644 --- a/packages/django-google-spanner/tests/unit/django_spanner/test_base.py +++ b/packages/django-google-spanner/tests/unit/django_spanner/test_base.py @@ -15,6 +15,8 @@ def setUp(self): import django_spanner.base django_spanner.base._SPANNER_CLIENT_CACHE = None + self.db_wrapper.connection = None + self.db_wrapper.settings_dict = dict(self.settings_dict) self.client_patcher = mock.patch("django_spanner.base.spanner.Client") self.mock_client = self.client_patcher.start() self.mock_client.return_value.instance.return_value.instance_id = ( @@ -44,15 +46,19 @@ def test_get_connection_params(self): self.assertEqual(params["option"], self.OPTIONS["option"]) def test_get_new_connection(self): - self.db_wrapper.Database = mock_database = mock.MagicMock() - mock_database.connect = mock_connection = mock.MagicMock() - conn_params = {"test_param": "dummy"} - self.db_wrapper.get_new_connection(conn_params) - mock_connection.assert_called_once_with( - self.INSTANCE_ID, - client=mock_database.connect.call_args[1]["client"], - **conn_params, - ) + orig_database = self.db_wrapper.Database + try: + self.db_wrapper.Database = mock_database = mock.MagicMock() + mock_database.connect = mock_connection = mock.MagicMock() + conn_params = {"test_param": "dummy"} + self.db_wrapper.get_new_connection(conn_params) + mock_connection.assert_called_once_with( + self.INSTANCE_ID, + client=mock_database.connect.call_args[1]["client"], + **conn_params, + ) + finally: + self.db_wrapper.Database = orig_database def test_init_connection_state(self): from google.cloud.spanner_dbapi.connection import Connection @@ -87,15 +93,28 @@ def test_is_usable(self): mock_connection.is_closed = False self.assertTrue(self.db_wrapper.is_usable()) - def test_is_usable_with_error(self): - from google.cloud.spanner_dbapi.exceptions import Error + def test_allow_transactions_in_auto_commit(self): + self.assertTrue(self.db_wrapper.allow_transactions_in_auto_commit) + self.db_wrapper.settings_dict["ALLOW_TRANSACTIONS_IN_AUTO_COMMIT"] = False + self.assertFalse(self.db_wrapper.allow_transactions_in_auto_commit) - self.db_wrapper.connection = mock_connection = mock.MagicMock() - mock_connection.cursor = mock.MagicMock(side_effect=Error) + def test_is_usable_with_error(self): + mock_cursor = mock.MagicMock() + mock_cursor.execute.side_effect = self.db_wrapper.Database.Error("error") + mock_connection = mock.MagicMock(is_closed=False) + mock_connection.cursor.return_value = mock_cursor + self.db_wrapper.connection = mock_connection self.assertFalse(self.db_wrapper.is_usable()) def test_start_transaction_under_autocommit(self): + mock_cursor = mock.MagicMock() self.db_wrapper.connection = mock_connection = mock.MagicMock() - mock_connection.cursor = mock_cursor = mock.MagicMock() + mock_connection.cursor.return_value = mock_cursor + self.db_wrapper._start_transaction_under_autocommit() - mock_cursor.assert_called_once_with() + mock_cursor.execute.assert_called_once_with("BEGIN") + + mock_cursor.reset_mock() + self.db_wrapper.settings_dict["ALLOW_TRANSACTIONS_IN_AUTO_COMMIT"] = False + self.db_wrapper._start_transaction_under_autocommit() + mock_cursor.execute.assert_called_once_with("SELECT 1") diff --git a/packages/django-google-spanner/tests/unit/django_spanner/test_compiler.py b/packages/django-google-spanner/tests/unit/django_spanner/test_compiler.py index 328dee8b1fbf..5ae70bd6cce1 100644 --- a/packages/django-google-spanner/tests/unit/django_spanner/test_compiler.py +++ b/packages/django-google-spanner/tests/unit/django_spanner/test_compiler.py @@ -4,6 +4,7 @@ # license that can be found in the LICENSE file or at # https://developers.google.com/open-source/licenses/bsd +from unittest import mock from django.core.exceptions import EmptyResultSet from django.db.models.query import QuerySet from django.db.utils import DatabaseError @@ -129,12 +130,12 @@ def test_get_combinator_sql_union_and_difference_query_together(self): self.assertEqual( sql_compiled, [ - "SELECT tests_number.num AS num FROM tests_number WHERE " - + "tests_number.num <= %s UNION DISTINCT SELECT * FROM (" + "(SELECT tests_number.num AS num FROM tests_number WHERE " + + "tests_number.num <= %s) UNION DISTINCT (SELECT * FROM (" + "SELECT tests_number.num AS num FROM tests_number WHERE " + "tests_number.num >= %s EXCEPT DISTINCT " + "SELECT tests_number.num AS num FROM tests_number " - + "WHERE tests_number.num = %s)" + + "WHERE tests_number.num = %s))" ], ) self.assertEqual(params, [1, 8, 10]) @@ -173,3 +174,53 @@ def test_get_combinator_sql_empty_queryset_raises_exception(self): compiler = SQLCompiler(QuerySet().query, self.connection, "default") with self.assertRaises(EmptyResultSet): compiler.get_combinator_sql("union", False) + + def test_get_combinator_sql_sliced_and_features(self): + qs1 = Number.objects.filter(num__lte=1) + qs2 = Number.objects.none() # Empty result set + qs3 = Number.objects.filter(num__gte=8) + qs4 = qs1.union(qs3) + + compiler = SQLCompiler(qs4.query, self.connection, "default") + compiler.connection.features.supports_slicing_ordering_in_compound = True + compiler.query.high_mark = 10 + sql, params = compiler.get_combinator_sql("union", False) + self.assertTrue(len(sql) > 0) + + # Test values_select set_values + qs_val = Number.objects.values("num") + qs_comb = qs_val.union(Number.objects.all()) + comp_val = SQLCompiler(qs_comb.query, self.connection, "default") + sql2, params2 = comp_val.get_combinator_sql("union", True) + self.assertTrue(len(sql2) > 0) + + def test_get_combinator_sql_edge_cases(self): + qs1, qs_empty = Number.objects.filter(num=1), Number.objects.none() + + # Test valid combinators (union, difference, subquery) + for comp in [ + SQLCompiler(qs1.union(qs_empty).query, self.connection, "default"), + SQLCompiler(qs1.difference(qs_empty).query, self.connection, "default"), + ]: + comp.query.subquery = True + self.assertTrue(len(comp.get_combinator_sql("union", False)[0]) > 0) + + # Test EmptyResultSet exceptions (intersection on empty, union on empty parts) + qs_none1, qs_none2 = Number.objects.none(), Number.objects.none() + for comp in [ + SQLCompiler(qs_empty.intersection(qs1).query, self.connection, "default"), + SQLCompiler(qs_none1.union(qs_none2).query, self.connection, "default"), + ]: + with self.assertRaises(EmptyResultSet): + comp.get_combinator_sql("union", False) + + # supports_parentheses_in_compound False with subquery combinator + comp_sub = SQLCompiler(qs1.union(Number.objects.filter(num=2)).query, self.connection, "default") + mock_sub = mock.MagicMock(as_sql=mock.MagicMock(return_value=("SELECT 1", []))) + mock_sub.query.combinator = "union" + mock_sub.query.is_sliced = False + mock_sub.get_order_by.return_value = [] + with mock.patch.object(self.connection.features, "supports_parentheses_in_compound", False): + with mock.patch.object(comp_sub.query.combined_queries[0], "get_compiler", return_value=mock_sub): + with mock.patch.object(comp_sub.query.combined_queries[1], "get_compiler", return_value=mock_sub): + self.assertIn("SELECT * FROM", comp_sub.get_combinator_sql("union", False)[0][0]) diff --git a/packages/django-google-spanner/tests/unit/django_spanner/test_creation.py b/packages/django-google-spanner/tests/unit/django_spanner/test_creation.py new file mode 100644 index 000000000000..5a1bfcccd2d6 --- /dev/null +++ b/packages/django-google-spanner/tests/unit/django_spanner/test_creation.py @@ -0,0 +1,97 @@ +# Copyright 2026 Google LLC +# +# Use of this source code is governed by a BSD-style +# license that can be found in the LICENSE file or at +# https://developers.google.com/open-source/licenses/bsd + +import os +from unittest import mock +from django_spanner.creation import DatabaseCreation +from tests.unit.django_spanner.simple_test import SpannerSimpleTestClass + + +class TestCreation(SpannerSimpleTestClass): + def setUp(self): + super().setUp() + self.db_wrapper.settings_dict = dict(self.settings_dict) + self.db_wrapper.settings_dict["TEST"] = {"NAME": "test_db"} + self.creation = DatabaseCreation(self.db_wrapper) + + def test_mark_skips(self): + with mock.patch("django.conf.settings.INSTALLED_APPS", ["django.contrib.contenttypes"]): + self.db_wrapper.features.skip_tests = ( + "django.contrib.contenttypes.models.ContentType", + ) + self.creation.mark_skips() + + def test_create_test_db_user_cancel(self): + with mock.patch.object(self.creation, "_execute_create_test_db", side_effect=Exception("err")): + with mock.patch("builtins.input", return_value="no"): + with self.assertRaises(SystemExit) as cm: + self.creation._create_test_db(verbosity=1, autoclobber=False, keepdb=False) + self.assertEqual(cm.exception.code, 1) + + def test_create_test_db_recreate_error_exit(self): + with mock.patch.object(self.creation, "_execute_create_test_db", side_effect=Exception("err")): + with mock.patch.object(self.creation, "_destroy_test_db", side_effect=Exception("destroy_err")): + with self.assertRaises(SystemExit) as cm: + self.creation._create_test_db(verbosity=1, autoclobber=True, keepdb=False) + self.assertEqual(cm.exception.code, 2) + + def test_create_test_db(self): + with mock.patch.dict(os.environ, {"RUNNING_SPANNER_BACKEND_TESTS": "1"}): + with mock.patch.object(self.creation, "mark_skips") as mock_mark: + with mock.patch("django.db.backends.base.creation.BaseDatabaseCreation.create_test_db"): + self.creation.create_test_db() + mock_mark.assert_called_once() + + # Without env var + env = dict(os.environ) + env.pop("RUNNING_SPANNER_BACKEND_TESTS", None) + with mock.patch.dict(os.environ, env, clear=True): + with mock.patch("django.db.backends.base.creation.BaseDatabaseCreation.create_test_db"): + self.creation.create_test_db() + + def test_mark_skips_success_and_attribute_error(self): + class DummyTestCase: + def test_foo(self): pass + with mock.patch("django.conf.settings.INSTALLED_APPS", ["tests"]): + with mock.patch("django_spanner.creation.import_string", return_value=DummyTestCase): + with mock.patch.object(self.creation.connection.features, "skip_tests", {"tests.DummyTestCase.test_foo", "tests.DummyTestCase.test_non_existent"}): + self.creation.mark_skips() + self.assertTrue(hasattr(DummyTestCase.test_foo, "__unittest_skip__")) + + def test_create_test_db_internal(self): + with mock.patch.object(self.creation, "_execute_create_test_db", side_effect=[Exception("db exists"), None]): + with mock.patch.object(self.creation, "_destroy_test_db") as mock_destroy: + with mock.patch.object(self.creation, "log") as mock_log: + self.creation._create_test_db(verbosity=1, autoclobber=True, keepdb=False) + mock_destroy.assert_called_once() + mock_log.assert_called() + + with mock.patch.object(self.creation, "_execute_create_test_db", side_effect=[Exception("db exists"), None]): + with mock.patch.object(self.creation, "_destroy_test_db"): + self.creation._create_test_db(verbosity=0, autoclobber=True, keepdb=False) + + def test_create_test_db_internal_error_keepdb(self): + with mock.patch.object(self.creation, "_execute_create_test_db", side_effect=Exception("create error")): + dbname = self.creation._create_test_db(verbosity=1, autoclobber=True, keepdb=True) + self.assertEqual(dbname, "test_db") + + def test_create_test_db_internal_error_recreate(self): + with mock.patch.object(self.creation, "_execute_create_test_db", side_effect=[Exception("err1"), None]): + with mock.patch.object(self.creation, "_destroy_test_db") as mock_destroy: + dbname = self.creation._create_test_db(verbosity=1, autoclobber=True, keepdb=False) + mock_destroy.assert_called_once() + self.assertEqual(dbname, "test_db") + + def test_execute_create_and_destroy_test_db(self): + mock_instance = mock.MagicMock() + with mock.patch.object(type(self.db_wrapper), "instance", new_callable=mock.PropertyMock(return_value=mock_instance)): + self.creation._execute_create_test_db(None, {"dbname": "test_db"}, keepdb=False) + mock_instance.database.assert_called_with("test_db") + mock_instance.database.return_value.create.assert_called_once() + + self.creation._destroy_test_db("test_db", verbosity=1) + mock_instance.database.assert_called_with("test_db") + mock_instance.database.return_value.drop.assert_called_once() diff --git a/packages/django-google-spanner/tests/unit/django_spanner/test_features.py b/packages/django-google-spanner/tests/unit/django_spanner/test_features.py new file mode 100644 index 000000000000..aff91afc48dd --- /dev/null +++ b/packages/django-google-spanner/tests/unit/django_spanner/test_features.py @@ -0,0 +1,27 @@ +# Copyright 2026 Google LLC +# +# Use of this source code is governed by a BSD-style +# license that can be found in the LICENSE file or at +# https://developers.google.com/open-source/licenses/bsd + +import importlib +import os +from unittest import mock + +from tests.unit.django_spanner.simple_test import SpannerSimpleTestClass + + +class TestFeatures(SpannerSimpleTestClass): + def test_features_emulator_and_env(self): + import django_spanner + import django_spanner.features + + with mock.patch.dict(os.environ, {"RUNNING_SPANNER_BACKEND_TESTS": "1", "SPANNER_EMULATOR_HOST": "localhost:9010"}): + with mock.patch.object(django_spanner, "USE_EMULATOR", True): + importlib.reload(django_spanner.features) + feat = django_spanner.features.DatabaseFeatures(self.db_wrapper) + self.assertFalse(feat.supports_foreign_keys) + self.assertFalse(feat.supports_json_field) + self.assertTrue(any("test_loaddata" in test for test in feat.skip_tests)) + + importlib.reload(django_spanner.features) diff --git a/packages/django-google-spanner/tests/unit/django_spanner/test_functions.py b/packages/django-google-spanner/tests/unit/django_spanner/test_functions.py index f5485b03b366..0ac2936d1c04 100644 --- a/packages/django-google-spanner/tests/unit/django_spanner/test_functions.py +++ b/packages/django-google-spanner/tests/unit/django_spanner/test_functions.py @@ -234,3 +234,18 @@ def test_substr(self): + "name_prefix FROM tests_author", ) self.assertEqual(params, (1, 5)) + + def test_chr(self): + from django.db.models.functions import Chr + q1 = Author.objects.values("num").annotate(chr_val=Chr("num")) + compiler = SQLCompiler(q1.query, self.connection, "default") + sql_query, params = compiler.query.as_sql(compiler, self.connection) + self.assertIn("CODE_POINTS_TO_STRING", sql_query) + + def test_json_array(self): + from django.db.models.functions import JSONObject + from django_spanner.functions import JSONArray + q1 = Author.objects.values("num").annotate(json_arr=JSONArray("num")) + compiler = SQLCompiler(q1.query, self.connection, "default") + sql_query, params = compiler.query.as_sql(compiler, self.connection) + self.assertIn("TO_JSON_STRING", sql_query) diff --git a/packages/django-google-spanner/tests/unit/django_spanner/test_init.py b/packages/django-google-spanner/tests/unit/django_spanner/test_init.py new file mode 100644 index 000000000000..8a3a2abffa43 --- /dev/null +++ b/packages/django-google-spanner/tests/unit/django_spanner/test_init.py @@ -0,0 +1,72 @@ +# Copyright 2026 Google LLC +# +# Use of this source code is governed by a BSD-style +# license that can be found in the LICENSE file or at +# https://developers.google.com/open-source/licenses/bsd + +import datetime +from unittest import mock +from django.db import DEFAULT_DB_ALIAS +from django.db.models import JSONField, AutoField +from google.api_core.datetime_helpers import DatetimeWithNanoseconds +from google.cloud.spanner_v1 import JsonObject + +from django_spanner import autofield_init +from tests.unit.django_spanner.simple_test import SpannerSimpleTestClass + + +class TestInit(SpannerSimpleTestClass): + def test_jsonfield_get_prep_value(self): + json_field = JSONField() + + # Dict value should be wrapped in JsonObject + res_dict = json_field.get_prep_value({"a": 1}) + self.assertIsInstance(res_dict, JsonObject) + self.assertEqual(res_dict, JsonObject({"a": 1})) + + # JsonObject value should be returned as-is + jo = JsonObject({"b": 2}) + self.assertEqual(json_field.get_prep_value(jo), jo) + + # Other types should be returned as-is + self.assertEqual(json_field.get_prep_value("string"), "string") + self.assertIsNone(json_field.get_prep_value(None)) + + def test_datetimewithnanoseconds_eq(self): + UTC = datetime.timezone.utc + dt = datetime.datetime(2020, 1, 10, 2, 44, 57, 999, UTC) + dt_diff = datetime.datetime(2021, 1, 10, 2, 44, 57, 999, UTC) + dtns1 = DatetimeWithNanoseconds(2020, 1, 10, 2, 44, 57, 999, UTC) + dtns2 = DatetimeWithNanoseconds(2020, 1, 10, 2, 44, 57, 999, UTC) + dtns3 = DatetimeWithNanoseconds(2020, 1, 10, 2, 44, 58, 000, UTC) + + # Equals another DatetimeWithNanoseconds (same instance or values) + self.assertTrue(dtns1 == dtns2) + self.assertFalse(dtns1 == dtns3) + self.assertTrue(dtns1 != dtns3) + self.assertFalse(dtns1 != dtns2) + + # Equals datetime.datetime with same ctime + self.assertTrue(dtns1 == dt) + self.assertFalse(dtns1 == dt_diff) + + # When old_datetimewithnanoseconds_eq is None + with mock.patch("django_spanner.old_datetimewithnanoseconds_eq", None): + self.assertTrue(dtns1 == dt) + self.assertFalse(dtns1 == dt_diff) + self.assertFalse(dtns1 == "not_a_dt") + + from django_spanner import datetimewithnanoseconds_eq + self.assertFalse(datetimewithnanoseconds_eq(dtns1, dtns3)) + self.assertFalse(datetimewithnanoseconds_eq(dtns1, dt_diff)) + + def test_gen_rand_int64(self): + from django_spanner import gen_rand_int64 + val = gen_rand_int64() + self.assertGreaterEqual(val, 0) + self.assertLessEqual(val, 0x7FFFFFFFFFFFFFFF) + + def test_autofield_init_options(self): + field = AutoField() + autofield_init(field) + self.assertTrue(field.blank) diff --git a/packages/django-google-spanner/tests/unit/django_spanner/test_introspection.py b/packages/django-google-spanner/tests/unit/django_spanner/test_introspection.py index 7091a31b2a46..223f62bfc226 100644 --- a/packages/django-google-spanner/tests/unit/django_spanner/test_introspection.py +++ b/packages/django-google-spanner/tests/unit/django_spanner/test_introspection.py @@ -71,13 +71,16 @@ def test_get_table_description(self): cursor = mock.MagicMock() def description(*args, **kwargs): - return [["name", TypeCode.STRING], ["age", TypeCode.INT64]] + return [["name", TypeCode.STRING], ["bio", TypeCode.STRING], ["age", TypeCode.INT64]] def get_table_column_schema(*args, **kwargs): column_details = {} column_details["name"] = ColumnDetails( null_ok=False, spanner_type="STRING(10)" ) + column_details["bio"] = ColumnDetails( + null_ok=True, spanner_type="STRING(MAX)" + ) column_details["age"] = ColumnDetails(null_ok=True, spanner_type="INT64") return column_details @@ -86,33 +89,8 @@ def get_table_column_schema(*args, **kwargs): table_description = db_introspection.get_table_description( cursor=cursor, table_name="Table_1" ) - self.assertEqual( - table_description, - [ - FieldInfo( - name="name", - type_code=TypeCode.STRING, - display_size=None, - internal_size=10, - precision=None, - scale=None, - null_ok=False, - default=None, - collation=None, - ), - FieldInfo( - name="age", - type_code=TypeCode.INT64, - display_size=None, - internal_size=None, - precision=None, - scale=None, - null_ok=True, - default=None, - collation=None, - ), - ], - ) + self.assertEqual(len(table_description), 3) + self.assertEqual(table_description[1].internal_size, "MAX") def test_get_primary_key_column(self): """ @@ -159,17 +137,19 @@ def test_get_constraints(self): cursor = mock.MagicMock() def run_sql_in_snapshot(*args, **kwargs): - # returns dummy data for 'CONSTRAINT_NAME, COLUMN_NAME' query. + # returns dummy data for 'CONSTRAINT_NAME, COLUMN_NAME' query with multi-column constraint if "CONSTRAINT_NAME, COLUMN_NAME" in args[0]: - return [["pk_constraint", "id"], ["name_constraint", "name"]] + return [["pk_constraint", "id"], ["pk_constraint", "sub_id"], ["name_constraint", "name"]] # returns dummy data for 'CONSTRAINT_NAME, CONSTRAINT_TYPE' query. if "CONSTRAINT_NAME, CONSTRAINT_TYPE" in args[0]: return [ ["pk_constraint", "PRIMARY KEY"], - ["FOREIGN KEY", "dept_id"], + ["name_constraint", "FOREIGN KEY"], + ["unadded_fk", "FOREIGN KEY"], + ["new_const", "CHECK"], ] - # returns dummy data for 'INFORMATION_SCHEMA.INDEXES' table query. - return [["pk_index", "id", "ASCENDING", "PRIMARY_KEY", True]] + # returns dummy data for 'INFORMATION_SCHEMA.INDEXES' table query with multi-column index. + return [["pk_index", "id", "ASCENDING", "PRIMARY_KEY", True], ["pk_index", "sub_id", "ASCENDING", "PRIMARY_KEY", True]] cursor.run_sql_in_snapshot = run_sql_in_snapshot constraints = db_introspection.get_constraints( @@ -181,7 +161,7 @@ def run_sql_in_snapshot(*args, **kwargs): { "pk_constraint": { "check": False, - "columns": ["id"], + "columns": ["id", "sub_id"], "foreign_key": None, "index": False, "orders": [], @@ -189,18 +169,8 @@ def run_sql_in_snapshot(*args, **kwargs): "type": None, "unique": True, }, - "name_constraint": { - "check": False, - "columns": ["name"], - "foreign_key": None, - "index": False, - "orders": [], - "primary_key": False, - "type": None, - "unique": False, - }, - "FOREIGN KEY": { - "check": False, + "new_const": { + "check": True, "columns": [], "foreign_key": None, "index": False, @@ -211,13 +181,40 @@ def run_sql_in_snapshot(*args, **kwargs): }, "pk_index": { "check": False, - "columns": ["id"], + "columns": ["id", "sub_id"], "foreign_key": None, "index": True, - "orders": ["ASCENDING"], + "orders": ["ASCENDING", "ASCENDING"], "primary_key": True, "type": "PRIMARY_KEY", "unique": True, }, }, ) + + def test_get_table_list_with_view(self): + db_introspection = DatabaseIntrospection(self.connection) + cursor = mock.MagicMock() + cursor.run_sql_in_snapshot.return_value = [["View_1", "VIEW"]] + table_list = db_introspection.get_table_list(cursor=cursor) + self.assertEqual(table_list, [TableInfo(name="View_1", type="v")]) + + def test_get_relations(self): + db_introspection = DatabaseIntrospection(self.connection) + cursor = mock.MagicMock() + cursor.run_sql_in_snapshot.return_value = [("author_id", "id", "author")] + relations = db_introspection.get_relations(cursor=cursor, table_name="book") + self.assertEqual(relations, {"author_id": ("id", "author")}) + + def test_get_key_columns(self): + db_introspection = DatabaseIntrospection(self.connection) + cursor = mock.MagicMock() + cursor.fetchall.return_value = [("author_id", "author", "id")] + keys = db_introspection.get_key_columns(cursor=cursor, table_name="book") + self.assertEqual(keys, [("author_id", "author", "id")]) + + def test_get_sequences(self): + db_introspection = DatabaseIntrospection(self.connection) + cursor = mock.MagicMock() + with self.assertRaises(NotImplementedError): + db_introspection.get_sequences(cursor=cursor, table_name="book") diff --git a/packages/django-google-spanner/tests/unit/django_spanner/test_lookups.py b/packages/django-google-spanner/tests/unit/django_spanner/test_lookups.py index deab4191bcc3..ec85c66ba699 100644 --- a/packages/django-google-spanner/tests/unit/django_spanner/test_lookups.py +++ b/packages/django-google-spanner/tests/unit/django_spanner/test_lookups.py @@ -5,6 +5,7 @@ # https://developers.google.com/open-source/licenses/bsd from decimal import Decimal +from unittest import mock from django.db.models import F @@ -271,3 +272,24 @@ def test_iexact_sql_query_case_insensitive_value_match(self): ) self.assertEqual(sql_compiled, expected_sql) self.assertEqual(params, ("abc",)) + + def test_in_lookup_str_param_cast(self): + from django.db.models.lookups import Exact + lookup = Exact(Author._meta.get_field("name").get_col("tests_author"), "10") + compiler = SQLCompiler(Author.objects.all().query, self.connection, "default") + field_mock = mock.MagicMock() + field_mock.rel_db_type.return_value = "INT64" + target_info = mock.MagicMock() + target_info.target_fields = [field_mock] + with mock.patch.object(lookup, "as_sql", return_value=("name = %s", ["10"])): + with mock.patch.object(lookup.lhs.output_field, "get_path_info", return_value=[target_info], create=True): + sql, params = Exact.as_spanner(lookup, compiler, self.connection) + self.assertEqual(params[0], 10) + + def test_regex_with_placeholder_and_params(self): + from django.db.models import Value + from django.db.models.functions import Concat + qs1 = Author.objects.filter(name__iexact=Concat(F("last_name"), Value("test"))).values("name") + compiler = SQLCompiler(qs1.query, self.connection, "default") + sql, params = compiler.as_sql() + self.assertTrue(len(sql) > 0) diff --git a/packages/django-google-spanner/tests/unit/django_spanner/test_operations.py b/packages/django-google-spanner/tests/unit/django_spanner/test_operations.py index bbadf5129e22..7086df8b9dcd 100644 --- a/packages/django-google-spanner/tests/unit/django_spanner/test_operations.py +++ b/packages/django-google-spanner/tests/unit/django_spanner/test_operations.py @@ -4,6 +4,7 @@ # license that can be found in the LICENSE file or at # https://developers.google.com/open-source/licenses/bsd +import os import uuid from base64 import b64encode from datetime import timedelta @@ -14,12 +15,151 @@ from django.db.utils import DatabaseError from google.cloud.spanner_dbapi.types import DateStr +from unittest import mock + from tests.unit.django_spanner.simple_test import SpannerSimpleTestClass class TestOperations(SpannerSimpleTestClass): - def test_max_name_length(self): - self.assertEqual(self.db_operations.max_name_length(), 128) + def test_execute_sql_flush(self): + cursor = mock.MagicMock() + cursor_cm = mock.MagicMock() + cursor_cm.__enter__.return_value = cursor + with mock.patch.object(self.db_operations.connection, "get_autocommit", return_value=True): + with mock.patch.object(self.db_operations.connection, "cursor", return_value=cursor_cm): + self.db_operations.execute_sql_flush(["DELETE FROM T1 WHERE 1=1"]) + cursor.execute.assert_called_once_with("DELETE FROM T1 WHERE 1=1") + + def test_execute_sql_flush_empty_and_autocommit_false_and_error(self): + # Empty list + self.assertIsNone(self.db_operations.execute_sql_flush([])) + + cursor = mock.MagicMock() + cursor_cm = mock.MagicMock() + cursor_cm.__enter__.return_value = cursor + with mock.patch.object(self.db_operations.connection, "get_autocommit", return_value=False): + with mock.patch.object(self.db_operations.connection, "set_autocommit") as mock_set_auto: + with mock.patch.object(self.db_operations.connection, "cursor", return_value=cursor_cm): + self.db_operations.execute_sql_flush(["DELETE FROM T1 WHERE 1=1"]) + mock_set_auto.assert_has_calls([mock.call(True), mock.call(False)]) + + # Exception during execution causes no-progress error + cursor_err = mock.MagicMock() + cursor_err.execute.side_effect = DatabaseError("db error") + cursor_err_cm = mock.MagicMock() + cursor_err_cm.__enter__.return_value = cursor_err + with mock.patch.object(self.db_operations.connection, "get_autocommit", return_value=True): + with mock.patch.object(self.db_operations.connection, "cursor", return_value=cursor_err_cm): + with self.assertRaises(DatabaseError): + self.db_operations.execute_sql_flush(["DELETE FROM T1"]) + + def test_execute_sql_flush_max_passes(self): + cursor = mock.MagicMock() + # 11 queries: 1 succeeds per pass for 10 passes, pass 11 triggers max_passes + side_effects = [] + for i in range(10): + side_effects.append(None) # first statement succeeds + side_effects.extend([DatabaseError("err")] * (10 - i)) # remaining fail + cursor.execute.side_effect = side_effects + cursor_cm = mock.MagicMock() + cursor_cm.__enter__.return_value = cursor + queries = ["DELETE FROM T%d" % i for i in range(11)] + with mock.patch.object(self.db_operations.connection, "get_autocommit", return_value=True): + with mock.patch.object(self.db_operations.connection, "cursor", return_value=cursor_cm): + with self.assertRaises(DatabaseError): + self.db_operations.execute_sql_flush(queries) + + def test_date_and_datetime_trunc_week(self): + sql_date, _ = self.db_operations.date_trunc_sql("week", "col", []) + self.assertIn("DATE_SUB", sql_date) + self.assertIn("DATE_ADD", sql_date) + + sql_date_day, _ = self.db_operations.date_trunc_sql("day", "col", []) + self.assertNotIn("DATE_SUB", sql_date_day) + + sql_dt, _ = self.db_operations.datetime_trunc_sql("week", "col", []) + self.assertIn("TIMESTAMP_SUB", sql_dt) + self.assertIn("TIMESTAMP_ADD", sql_dt) + + sql_dt_day, _ = self.db_operations.datetime_trunc_sql("day", "col", []) + self.assertNotIn("TIMESTAMP_SUB", sql_dt_day) + + def test_limit_offset_params_and_savepoints(self): + self.assertEqual(self.db_operations.integer_field_range("IntegerField"), (-9223372036854775808, 9223372036854775807)) + lim, off = self.db_operations._get_limit_offset_params(5, None) + self.assertEqual(off, 5) + self.assertEqual(lim, 9223372036854775807 - 5) + + lim0, off0 = self.db_operations._get_limit_offset_params(0, 10) + self.assertEqual(off0, 0) + self.assertEqual(lim0, 10) + + self.assertEqual(self.db_operations.savepoint_create_sql("s1"), "SELECT 1") + self.assertEqual(self.db_operations.savepoint_commit_sql("s1"), "SELECT 1") + self.assertEqual(self.db_operations.savepoint_rollback_sql("s1"), "SELECT 1") + + def test_adapt_and_convert_datetime_time_fields(self): + from django.db.models import F, DateTimeField, TimeField, BinaryField, UUIDField + from google.api_core.datetime_helpers import DatetimeWithNanoseconds + import datetime + + # adapt_datetimefield_value with resolve_expression + expr = F("created") + self.assertEqual(self.db_operations.adapt_datetimefield_value(expr), expr) + self.assertIsNone(self.db_operations.adapt_datetimefield_value(None)) + + # adapt_timefield_value + self.assertEqual(self.db_operations.adapt_timefield_value(expr), expr) + self.assertIsNone(self.db_operations.adapt_timefield_value(None)) + t_val = datetime.time(12, 30, 45) + res_t = self.db_operations.adapt_timefield_value(t_val) + self.assertEqual(res_t, "0001-01-01T12:30:45.000000Z") + + # get_db_converters + for f_cls in [DateTimeField, TimeField, BinaryField, UUIDField]: + mock_expr = mock.MagicMock() + mock_expr.output_field.get_internal_type.return_value = f_cls().__class__.__name__ + convs = self.db_operations.get_db_converters(mock_expr) + self.assertTrue(len(convs) > 0) + + # convert_datetimefield_value & convert_timefield_value + dtns = DatetimeWithNanoseconds(2020, 1, 1, 12, 0, 0) + self.assertIsNone(self.db_operations.convert_datetimefield_value(None, None, None)) + with mock.patch("django.conf.settings.USE_TZ", False): + conv_dt = self.db_operations.convert_datetimefield_value(dtns, None, None) + self.assertEqual(conv_dt.year, 2020) + + self.assertIsNone(self.db_operations.convert_timefield_value(None, None, None)) + conv_t = self.db_operations.convert_timefield_value(dtns, None, None) + self.assertEqual(conv_t.hour, 12) + + def test_quote_name_env_and_bulk_batch_size(self): + with mock.patch.dict(os.environ, {"RUNNING_SPANNER_BACKEND_TESTS": "1"}): + self.assertEqual(self.db_operations.quote_name("my name"), "my_name") + self.assertEqual(self.db_operations.bulk_batch_size(["f1", "f2"], []), 900 // 2) + + def test_adapt_datetimefield_value_value_error(self): + import datetime + from django.utils import timezone + dt_aware = timezone.make_aware(datetime.datetime(2020, 1, 1, 12, 0, 0), datetime.timezone.utc) + with mock.patch("django.conf.settings.USE_TZ", False): + with self.assertRaises(ValueError): + self.db_operations.adapt_datetimefield_value(dt_aware) + + # Aware datetime when USE_TZ is True executes line 254 (make_naive) + self.db_operations.connection.settings_dict["TIME_ZONE"] = "UTC" + with mock.patch("django.conf.settings.USE_TZ", True): + res = self.db_operations.adapt_datetimefield_value(dt_aware) + self.assertIsNotNone(res) + + # Naive datetime takes 252->260 branch directly + dt_naive = datetime.datetime(2020, 1, 1, 12, 0, 0) + res_naive = self.db_operations.adapt_datetimefield_value(dt_naive) + self.assertIsNotNone(res_naive) + + def test_combine_expression_xor(self): + res = self.db_operations.combine_expression("#", ["a", "b"]) + self.assertIn("^", res) def test_quote_name(self): quoted_name = self.db_operations.quote_name("abc") @@ -110,159 +250,30 @@ def test_convert_uuidfield_value_none(self): ), ) - def test_date_extract_sql(self): - self.assertEqual( - self.db_operations.date_extract_sql("week", "dummy_field"), - ("EXTRACT(isoweek FROM dummy_field)", None), - ) - - def test_date_extract_sql_lookup_type_dayofweek(self): - self.assertEqual( - self.db_operations.date_extract_sql("dayofweek", "dummy_field"), - ("EXTRACT(dayofweek FROM dummy_field)", None), - ) - - def test_datetime_extract_sql(self): + def test_sql_expressions_and_conversions(self): + ops = self.db_operations + self.assertEqual(ops.date_extract_sql("week", "f"), ("EXTRACT(isoweek FROM f)", None)) + self.assertEqual(ops.date_extract_sql("dayofweek", "f"), ("EXTRACT(dayofweek FROM f)", None)) + self.assertEqual(ops.time_extract_sql("dayofweek", "f"), ('EXTRACT(dayofweek FROM f AT TIME ZONE "UTC")', None)) + self.assertEqual(ops.time_trunc_sql("dayofweek", "f", None), ('TIMESTAMP_TRUNC(f, dayofweek, "UTC")', None)) + self.assertEqual(ops.datetime_cast_date_sql("f", None, "IST"), ('DATE(f, "IST")', None)) + self.assertEqual(ops.date_interval_sql(timedelta(days=1)), "INTERVAL 86400000000 MICROSECOND") + self.assertEqual(ops.format_for_duration_arithmetic(1200), "INTERVAL 1200 MICROSECOND") + self.assertEqual(ops.combine_expression("%%", ["10", "2"]), "MOD(10, 2)") + self.assertEqual(ops.combine_expression("^", ["10", "2"]), "POWER(10, 2)") + self.assertEqual(ops.combine_expression(">>", ["10", "2"]), "CAST(FLOOR(10 / POW(2, 2)) AS INT64)") + self.assertEqual(ops.combine_expression("*", ["10", "2"]), "10 * 2") + self.assertEqual(ops.combine_duration_expression("+", ["t", "i"]), "TIMESTAMP_ADD(t, i)") + self.assertEqual(ops.combine_duration_expression("-", ["t", "i"]), "TIMESTAMP_SUB(t, i)") + self.assertEqual(ops.lookup_cast("contains"), "CAST(%s AS STRING)") + self.assertEqual(ops.lookup_cast("dummy"), "%s") + + with self.assertRaises(DatabaseError): + ops.combine_duration_expression("*", ["t", "i"]) + + for use_tz in [True, False]: + settings.USE_TZ = use_tz + tz = "IST" if use_tz else "UTC" + self.assertIn(tz, ops.datetime_extract_sql("dayofweek", "f", None, "IST")[0]) + self.assertIn(tz, ops.datetime_cast_time_sql("f", None, "IST")[0]) settings.USE_TZ = True - self.assertEqual( - self.db_operations.datetime_extract_sql( - "dayofweek", "dummy_field", None, "IST" - ), - ( - 'EXTRACT(dayofweek FROM dummy_field AT TIME ZONE "IST")', - None, - ), - ) - - def test_datetime_extract_sql_use_tz_false(self): - settings.USE_TZ = False - self.assertEqual( - self.db_operations.datetime_extract_sql( - "dayofweek", "dummy_field", None, "IST" - ), - ( - 'EXTRACT(dayofweek FROM dummy_field AT TIME ZONE "UTC")', - None, - ), - ) - settings.USE_TZ = True # reset changes. - - def test_time_extract_sql(self): - self.assertEqual( - self.db_operations.time_extract_sql("dayofweek", "dummy_field"), - ( - 'EXTRACT(dayofweek FROM dummy_field AT TIME ZONE "UTC")', - None, - ), - ) - - def test_time_trunc_sql(self): - self.assertEqual( - self.db_operations.time_trunc_sql("dayofweek", "dummy_field", None), - ('TIMESTAMP_TRUNC(dummy_field, dayofweek, "UTC")', None), - ) - - def test_datetime_cast_date_sql(self): - self.assertEqual( - self.db_operations.datetime_cast_date_sql("dummy_field", None, "IST"), - ('DATE(dummy_field, "IST")', None), - ) - - def test_datetime_cast_time_sql(self): - settings.USE_TZ = True - self.assertEqual( - self.db_operations.datetime_cast_time_sql("dummy_field", None, "IST"), - ( - "TIMESTAMP(FORMAT_TIMESTAMP('%Y-%m-%d %R:%E9S %Z', dummy_field, 'IST'))", - None, - ), - ) - - def test_datetime_cast_time_sql_use_tz_false(self): - settings.USE_TZ = False - self.assertEqual( - self.db_operations.datetime_cast_time_sql("dummy_field", None, "IST"), - ( - "TIMESTAMP(FORMAT_TIMESTAMP('%Y-%m-%d %R:%E9S %Z', dummy_field, 'UTC'))", - None, - ), - ) - settings.USE_TZ = True # reset changes. - - def test_date_interval_sql(self): - self.assertEqual( - self.db_operations.date_interval_sql(timedelta(days=1)), - "INTERVAL 86400000000 MICROSECOND", - ) - - def test_format_for_duration_arithmetic(self): - self.assertEqual( - self.db_operations.format_for_duration_arithmetic(1200), - "INTERVAL 1200 MICROSECOND", - ) - - def test_combine_expression_mod(self): - self.assertEqual( - self.db_operations.combine_expression("%%", ["10", "2"]), - "MOD(10, 2)", - ) - - def test_combine_expression_power(self): - self.assertEqual( - self.db_operations.combine_expression("^", ["10", "2"]), - "POWER(10, 2)", - ) - - def test_combine_expression_bit_extention(self): - self.assertEqual( - self.db_operations.combine_expression(">>", ["10", "2"]), - "CAST(FLOOR(10 / POW(2, 2)) AS INT64)", - ) - - def test_combine_expression_multiply(self): - self.assertEqual( - self.db_operations.combine_expression("*", ["10", "2"]), - "10 * 2", - ) - - def test_combine_duration_expression_add(self): - self.assertEqual( - self.db_operations.combine_duration_expression( - "+", - ['TIMESTAMP "2008-12-25 15:30:00+00', "INTERVAL 10 MINUTE"], - ), - 'TIMESTAMP_ADD(TIMESTAMP "2008-12-25 15:30:00+00, INTERVAL 10 MINUTE)', - ) - - def test_combine_duration_expression_subtract(self): - self.assertEqual( - self.db_operations.combine_duration_expression( - "-", - ['TIMESTAMP "2008-12-25 15:30:00+00', "INTERVAL 10 MINUTE"], - ), - 'TIMESTAMP_SUB(TIMESTAMP "2008-12-25 15:30:00+00, INTERVAL 10 MINUTE)', - ) - - def test_combine_duration_expression_database_error(self): - msg = "Invalid connector for timedelta:" - with self.assertRaisesRegex(DatabaseError, msg): - self.db_operations.combine_duration_expression( - "*", - ['TIMESTAMP "2008-12-25 15:30:00+00', "INTERVAL 10 MINUTE"], - ) - - def test_lookup_cast_match_lookup_type(self): - self.assertEqual( - self.db_operations.lookup_cast( - "contains", - ), - "CAST(%s AS STRING)", - ) - - def test_lookup_cast_unmatched_lookup_type(self): - self.assertEqual( - self.db_operations.lookup_cast( - "dummy", - ), - "%s", - ) diff --git a/packages/django-google-spanner/tests/unit/django_spanner/test_schema.py b/packages/django-google-spanner/tests/unit/django_spanner/test_schema.py index b7ef7cec39ec..c52c2de0d01a 100644 --- a/packages/django-google-spanner/tests/unit/django_spanner/test_schema.py +++ b/packages/django-google-spanner/tests/unit/django_spanner/test_schema.py @@ -47,6 +47,115 @@ def test_skip_default(self): schema_editor = DatabaseSchemaEditor(self.connection) self.assertTrue(schema_editor.skip_default(field=None)) + def test_create_model_foreign_key_and_check_constraint(self): + from django.db.models import Model, ForeignKey, IntegerField, CASCADE + class Book(Model): + author = ForeignKey(Author, on_delete=CASCADE) + pages = IntegerField() + class Meta: + app_label = "tests" + + with DatabaseSchemaEditor(self.connection) as schema_editor: + schema_editor.execute = mock.MagicMock() + + # Test sql_create_inline_fk + schema_editor.sql_create_inline_fk = "CONSTRAINT FK FOREIGN KEY (%(from_column_norm)s) REFERENCES %(to_table_norm)s (%(to_column_norm)s)" + schema_editor.create_model(Book) + + # Test supports_foreign_keys with no inline_fk (False and True) + schema_editor.sql_create_inline_fk = None + with mock.patch.object(schema_editor.connection.features, "supports_foreign_keys", False): + schema_editor.create_model(Book) + with mock.patch.object(schema_editor.connection.features, "supports_foreign_keys", True): + schema_editor.create_model(Book) + + # Test check constraint on column + f_pages = Book._meta.get_field("pages") + with mock.patch.object(f_pages, "db_parameters", return_value={"type": "INT64", "check": "pages > 0"}): + schema_editor.create_model(Book) + + def test_create_model_unique_together_and_constraints(self): + from django.db.models import Model, OneToOneField, IntegerField, UniqueConstraint, CheckConstraint, Q, CASCADE + class Profile(Model): + author = OneToOneField(Author, on_delete=CASCADE, primary_key=True) + code = IntegerField() + class Meta: + app_label = "tests" + unique_together = [("code",)] + constraints = [ + UniqueConstraint(fields=["code"], name="unique_code_idx"), + CheckConstraint(check=Q(code__gt=0), name="code_gt_zero"), + ] + + with DatabaseSchemaEditor(self.connection) as schema_editor: + schema_editor.execute = mock.MagicMock() + schema_editor.create_model(Profile) + + def test_m2m_field_schema_operations(self): + from django.db.models import Model, ManyToManyField, IntegerField + class Article(Model): + authors = ManyToManyField(Author) + num = IntegerField() + class Meta: + app_label = "tests" + + with DatabaseSchemaEditor(self.connection) as schema_editor: + schema_editor.execute = mock.MagicMock() + schema_editor._constraint_names = mock.MagicMock(return_value=[]) + + # Test col_type_suffix and empty tablespace_sql + with mock.patch.object(Article._meta.get_field("num"), "db_type_suffix", return_value="SUFFIX"): + with mock.patch.object(Article._meta, "db_tablespace", "tbl"): + with mock.patch.object(schema_editor.connection.ops, "tablespace_sql", return_value=""): + schema_editor.create_model(Article) + + m2m_field = Article._meta.get_field("authors") + schema_editor.add_field(Article, m2m_field) + schema_editor.remove_field(Article, m2m_field) + schema_editor.alter_field(Article, m2m_field, m2m_field) + + # Test add_field with FK, unique, and check constraint + from django.db.models import ForeignKey, CASCADE, IntegerField, CharField + fk_field = ForeignKey(Author, on_delete=CASCADE) + fk_field.set_attributes_from_name("author_fk") + fk_field.model = Author + fk_field.unique = True + with mock.patch.object(fk_field, "db_parameters", return_value={"type": "INT64", "check": "num > 0"}): + with mock.patch.object(schema_editor.connection.features, "supports_foreign_keys", True): + schema_editor.add_field(Author, fk_field) + + # Test M2M, None column SQL, and db_default edge cases + for m in [Article, Author]: + schema_editor.create_model(m) + self.assertIsNone(schema_editor.column_sql(Article, m2m_field)[0]) + + def_field = IntegerField(db_default=10) + def_field.set_attributes_from_name("def_num") + def_field.model = Author + self.assertIn("DEFAULT", schema_editor.column_sql(Author, def_field)[0]) + + with mock.patch.object(schema_editor, "db_default_sql", return_value=(None, [])): + self.assertNotIn("DEFAULT", schema_editor.column_sql(Author, def_field)[0]) + + with mock.patch.object(schema_editor, "column_sql", return_value=(None, None)): + schema_editor.create_model(Article) + + # Test _alter_column_type_sql with same type, empty_strings_allowed, tablespace + char_field1 = CharField(max_length=50, db_tablespace="tbl", unique=True) + char_field1.set_attributes_from_name("title") + char_field1.model = Author + char_field2 = CharField(max_length=100, db_tablespace="tbl", unique=True) + char_field2.set_attributes_from_name("title") + char_field2.model = Author + with mock.patch.object(schema_editor.connection.features, "interprets_empty_strings_as_nulls", True): + with mock.patch.object(schema_editor.connection.features, "supports_tablespaces", True): + schema_editor._alter_column_type_sql(Author, char_field1, char_field2, "STRING(100)") + # Test column_sql with tablespace & empty string null (lines 371, 381) + schema_editor.column_sql(Author, char_field1) + # Same type returns sql statement + sql_same, _ = schema_editor._alter_column_type_sql(Author, char_field1, char_field1, "STRING(50)") + self.assertIsNotNone(sql_same) + def test_create_model(self): """ Tries creating a model's table. @@ -77,10 +186,11 @@ def test_delete_model(self): """ with DatabaseSchemaEditor(self.connection) as schema_editor: schema_editor.execute = mock.MagicMock() - schema_editor._constraint_names = mock.MagicMock() - schema_editor.delete_model(Author) + schema_editor._constraint_names = mock.MagicMock(return_value=[]) + with mock.patch.object(schema_editor.connection.features, "supports_foreign_keys", True): + schema_editor.delete_model(Author) - schema_editor.execute.assert_called_once_with( + schema_editor.execute.assert_called_with( "DROP TABLE tests_author", ) @@ -458,9 +568,51 @@ def test_autofield_spanner_as_non_default_db_random_generation_enabled( connections.settings["secondary"]["ENGINE"] = "django_spanner" del connections.settings["secondary"]["RANDOM_ID_GENERATION_ENABLED"] - def test_autofield_random_generation_disabled(self): - """Spanner, default is not provided.""" - connections.settings["default"]["RANDOM_ID_GENERATION_ENABLED"] = "false" - field = AutoField(name="field_name") - assert gen_rand_int64 != field.default - del connections.settings["default"]["RANDOM_ID_GENERATION_ENABLED"] + def test_schema_editor_utils_and_sql_formatting(self): + with DatabaseSchemaEditor(self.connection) as se: + se.execute = mock.MagicMock() + + # Quote values and prepare default + self.assertEqual(se.quote_value("it's"), "'it''s'") + self.assertEqual(se.quote_value(True), "TRUE") + self.assertEqual(se.prepare_default(123), "123") + + # Unique & check SQL + self.assertIsNone(se._unique_sql(Author, [Author._meta.get_field("name")], "idx")) + with mock.patch("django_spanner.schema.USE_EMULATOR", False): + self.assertIsNotNone(se._check_sql("chk", "num > 0")) + se.deferred_sql.clear() + se._unique_sql(Author, [Author._meta.get_field("name")], "idx") + + with mock.patch.object(se, "_create_unique_sql", return_value=None): + se._unique_sql(Author, [Author._meta.get_field("name")], "idx") + + # Column SQL generated & db_default + gen_f = IntegerField() + gen_f.generated = True + gen_f.set_attributes_from_name("g") + gen_f.model = Author + gen_f.generated_sql = lambda c: ("num * %s + %s + %s + %s", ["2", True, None, 5]) + self.assertIn("STORED", se.column_sql(Author, gen_f)[0]) + + def_f = IntegerField(db_default=10) + def_f.set_attributes_from_name("d") + def_f.model = Author + se.db_default_sql = lambda f: ("%s + %s + %s + %s", ["10", False, None, 3]) + self.assertIn("DEFAULT", se.column_sql(Author, def_f)[0]) + + # Skip default & column type alter + self.assertFalse(se.skip_default(gen_f)) + self.assertFalse(se.skip_default(def_f)) + + f_null = IntegerField(null=True) + f_null.set_attributes_from_name("v") + self.assertTrue(len(se._alter_column_type_sql(Author, f_null, f_null, "INT64")[0]) > 0) + + # None definition add_field & tablespace + with mock.patch.object(se, "column_sql", return_value=(None, None)): + self.assertIsNone(se.add_field(Author, gen_f)) + + with mock.patch.object(Author._meta, "db_tablespace", "tbl"): + with mock.patch.object(se.connection.ops, "tablespace_sql", return_value="IN tbl"): + se.create_model(Author)