From a56407035dcd8ace4a9670b6ee1dab1b662a7ca0 Mon Sep 17 00:00:00 2001 From: "petr.prikryl" Date: Wed, 9 Nov 2022 11:59:56 +0100 Subject: [PATCH 1/2] Django 4.1 support --- migrate_sql/apps.py | 23 +++ migrate_sql/autodetector.py | 78 ++++++---- migrate_sql/management/__init__.py | 0 migrate_sql/management/commands/__init__.py | 0 .../management/commands/makemigrations.py | 136 ------------------ 5 files changed, 75 insertions(+), 162 deletions(-) create mode 100644 migrate_sql/apps.py delete mode 100644 migrate_sql/management/__init__.py delete mode 100644 migrate_sql/management/commands/__init__.py delete mode 100644 migrate_sql/management/commands/makemigrations.py diff --git a/migrate_sql/apps.py b/migrate_sql/apps.py new file mode 100644 index 0000000..f1d8de2 --- /dev/null +++ b/migrate_sql/apps.py @@ -0,0 +1,23 @@ +from django.apps import AppConfig +from django.core.management.commands import makemigrations + +from migrate_sql.autodetector import MigrationAutodetectorMixin + + +def patch_autodetector(): + if not issubclass(makemigrations.MigrationAutodetector, MigrationAutodetectorMixin): + makemigrations.MigrationAutodetector = type( + "MigrationAutodetector", + ( + MigrationAutodetectorMixin, + makemigrations.MigrationAutodetector, + ), + {}, + ) + + +class MigrateSQLConfig(AppConfig): + name = "migrate_sql" + + def ready(self): + patch_autodetector() diff --git a/migrate_sql/autodetector.py b/migrate_sql/autodetector.py index 53613aa..cba24c0 100644 --- a/migrate_sql/autodetector.py +++ b/migrate_sql/autodetector.py @@ -1,14 +1,20 @@ -from django.db.migrations.autodetector import MigrationAutodetector as DjangoMigrationAutodetector from django.db.migrations.operations import RunSQL from django.utils.datastructures import OrderedSet -from migrate_sql.operations import (AlterSQL, ReverseAlterSQL, CreateSQL, DeleteSQL, AlterSQLState) -from migrate_sql.graph import SQLStateGraph +from migrate_sql.operations import ( + AlterSQL, + ReverseAlterSQL, + CreateSQL, + DeleteSQL, + AlterSQLState, +) +from migrate_sql.graph import SQLStateGraph, build_current_graph class SQLBlob(object): pass + # Dummy object used to identify django dependency as the one used by this tool only. SQL_BLOB = SQLBlob() @@ -59,34 +65,37 @@ def is_sql_equal(sqls1, sqls2): def get_ancestors(node): """Logic extracted from Django <2.2 as this is dropped in later versions.""" - if '_ancestors' not in node.__dict__: + if "_ancestors" not in node.__dict__: ancestors = [] for parent in sorted(node.parents, reverse=True): ancestors += get_ancestors(parent) ancestors.append(node.key) - node.__dict__['_ancestors'] = list(OrderedSet(ancestors)) - return node.__dict__['_ancestors'] + node.__dict__["_ancestors"] = list(OrderedSet(ancestors)) + return node.__dict__["_ancestors"] def get_descendants(node): """Logic extracted from Django <2.2 as this is dropped in later versions.""" - if '_descendants' not in node.__dict__: + if "_descendants" not in node.__dict__: descendants = [] for child in sorted(node.children, reverse=True): descendants += get_descendants(child) descendants.append(node.key) - node.__dict__['_descendants'] = list(OrderedSet(descendants)) - return node.__dict__['_descendants'] + node.__dict__["_descendants"] = list(OrderedSet(descendants)) + return node.__dict__["_descendants"] -class MigrationAutodetector(DjangoMigrationAutodetector): +class MigrationAutodetectorMixin: """ Substitutes Django's MigrationAutodetector class, injecting SQL migrations logic. """ - def __init__(self, from_state, to_state, questioner=None, to_sql_graph=None): - super(MigrationAutodetector, self).__init__(from_state, to_state, questioner) - self.to_sql_graph = to_sql_graph - self.from_sql_graph = getattr(self.from_state, 'sql_state', None) or SQLStateGraph() + + def __init__(self, *args, **kwargs): + super().__init__(*args, **kwargs) + self.to_sql_graph = build_current_graph() + self.from_sql_graph = ( + getattr(self.from_state, "sql_state", None) or SQLStateGraph() + ) self.from_sql_graph.build_graph() self._sql_operations = [] @@ -114,7 +123,9 @@ def assemble_changes(self, keys, resolve_keys, sql_state): sql_item = sql_state.nodes[key] ancs = get_ancestors(node)[:-1] ancs.reverse() - pos = next((i for i, k in enumerate(result_keys) if k in ancs), len(result_keys)) + pos = next( + (i for i, k in enumerate(result_keys) if k in ancs), len(result_keys) + ) result_keys.insert(pos, key) if key in resolve_keys and not sql_item.replace: @@ -132,7 +143,10 @@ def add_sql_operation(self, app_label, sql_name, operation, dependencies): Add SQL operation and register it to be used as dependency for further sequential operations. """ - deps = [(dp[0], SQL_BLOB, dp[1], self._sql_operations.get(dp)) for dp in dependencies] + deps = [ + (dp[0], SQL_BLOB, dp[1], self._sql_operations.get(dp)) + for dp in dependencies + ] self.add_operation(app_label, operation, dependencies=deps) self._sql_operations[(app_label, sql_name)] = operation @@ -147,11 +161,17 @@ def _generate_reversed_sql(self, keys, changed_keys): app_label, sql_name = key old_item = self.from_sql_graph.nodes[key] new_item = self.to_sql_graph.nodes[key] - if not old_item.reverse_sql or old_item.reverse_sql == RunSQL.noop or new_item.replace: + if ( + not old_item.reverse_sql + or old_item.reverse_sql == RunSQL.noop + or new_item.replace + ): continue # migrate backwards - operation = ReverseAlterSQL(sql_name, old_item.reverse_sql, reverse_sql=old_item.sql) + operation = ReverseAlterSQL( + sql_name, old_item.reverse_sql, reverse_sql=old_item.sql + ) sql_deps = [n.key for n in self.from_sql_graph.node_map[key].children] sql_deps.append(key) self.add_sql_operation(app_label, sql_name, operation, sql_deps) @@ -173,14 +193,15 @@ def _generate_sql(self, keys, changed_keys): # state_reverse_sql, the latter one will be used for building state forward # instead of reverse_sql. if new_item.replace: - kwargs['state_reverse_sql'] = reverse_sql + kwargs["state_reverse_sql"] = reverse_sql reverse_sql = self.from_sql_graph.nodes[key].sql else: operation_cls = CreateSQL - kwargs = {'dependencies': list(sql_deps)} + kwargs = {"dependencies": list(sql_deps)} operation = operation_cls( - sql_name, new_item.sql, reverse_sql=reverse_sql, **kwargs) + sql_name, new_item.sql, reverse_sql=reverse_sql, **kwargs + ) sql_deps.append(key) self.add_sql_operation(app_label, sql_name, operation, sql_deps) @@ -198,8 +219,11 @@ def _generate_altered_sql_dependencies(self, dep_changed_keys): """ for key, removed_deps, added_deps in dep_changed_keys: app_label, sql_name = key - operation = AlterSQLState(sql_name, add_dependencies=tuple(added_deps), - remove_dependencies=tuple(removed_deps)) + operation = AlterSQLState( + sql_name, + add_dependencies=tuple(added_deps), + remove_dependencies=tuple(removed_deps), + ) sql_deps = [key] self.add_sql_operation(app_label, sql_name, operation, sql_deps) @@ -210,7 +234,9 @@ def _generate_delete_sql(self, delete_keys): for key in delete_keys: app_label, sql_name = key old_node = self.from_sql_graph.nodes[key] - operation = DeleteSQL(sql_name, old_node.reverse_sql, reverse_sql=old_node.sql) + operation = DeleteSQL( + sql_name, old_node.reverse_sql, reverse_sql=old_node.sql + ) sql_deps = [n.key for n in self.from_sql_graph.node_map[key].children] sql_deps.append(key) self.add_sql_operation(app_label, sql_name, operation, sql_deps) @@ -263,7 +289,7 @@ def check_dependency(self, operation, dependency): # NOTE: we follow the sort order created by `assemble_changes` so we build a fixed chain # of operations. thus we should match exact operation here. return dependency[3] == operation - return super(MigrationAutodetector, self).check_dependency(operation, dependency) + return super().check_dependency(operation, dependency) def generate_altered_fields(self): """ @@ -271,6 +297,6 @@ def generate_altered_fields(self): divided into smaller methods/functions for easier enhancement and substitution. So far we're doing all the SQL magic in this method. """ - result = super(MigrationAutodetector, self).generate_altered_fields() + result = super().generate_altered_fields() self.generate_sql_changes() return result diff --git a/migrate_sql/management/__init__.py b/migrate_sql/management/__init__.py deleted file mode 100644 index e69de29..0000000 diff --git a/migrate_sql/management/commands/__init__.py b/migrate_sql/management/commands/__init__.py deleted file mode 100644 index e69de29..0000000 diff --git a/migrate_sql/management/commands/makemigrations.py b/migrate_sql/management/commands/makemigrations.py deleted file mode 100644 index d249f6f..0000000 --- a/migrate_sql/management/commands/makemigrations.py +++ /dev/null @@ -1,136 +0,0 @@ -""" -Replaces built-in Django command and forces it generate SQL item modification operations -into regular Django migrations. -""" - -import sys - -from django.core.management.commands.makemigrations import Command as MakeMigrationsCommand -from django.db.migrations.loader import MigrationLoader -from django.db.migrations import Migration -from django.core.management.base import CommandError, no_translations -from django.db.migrations.questioner import InteractiveMigrationQuestioner -from django.apps import apps -from django.db.migrations.state import ProjectState - -from migrate_sql.autodetector import MigrationAutodetector -from migrate_sql.graph import build_current_graph - - -class Command(MakeMigrationsCommand): - - @no_translations - def handle(self, *app_labels, **options): - - self.verbosity = options.get('verbosity') - self.interactive = options.get('interactive') - self.dry_run = options.get('dry_run', False) - self.merge = options.get('merge', False) - self.empty = options.get('empty', False) - self.migration_name = options.get('name', None) - self.exit_code = options.get('exit_code', False) - self.include_header = options.get('include_header', True) - check_changes = options.get('check_changes', False) - - # Make sure the app they asked for exists - app_labels = set(app_labels) - bad_app_labels = set() - for app_label in app_labels: - try: - apps.get_app_config(app_label) - except LookupError: - bad_app_labels.add(app_label) - if bad_app_labels: - for app_label in bad_app_labels: - self.stderr.write("App '%s' could not be found. Is it in INSTALLED_APPS?" % app_label) - sys.exit(2) - - # Load the current graph state. Pass in None for the connection so - # the loader doesn't try to resolve replaced migrations from DB. - loader = MigrationLoader(None, ignore_no_migrations=True) - - # Before anything else, see if there's conflicting apps and drop out - # hard if there are any and they don't want to merge - conflicts = loader.detect_conflicts() - - # If app_labels is specified, filter out conflicting migrations for unspecified apps - if app_labels: - conflicts = { - app_label: conflict for app_label, conflict in conflicts.items() - if app_label in app_labels - } - - if conflicts and not self.merge: - name_str = "; ".join( - "%s in %s" % (", ".join(names), app) - for app, names in conflicts.items() - ) - raise CommandError( - "Conflicting migrations detected (%s).\nTo fix them run " - "'python manage.py makemigrations --merge'" % name_str - ) - - # If they want to merge and there's nothing to merge, then politely exit - if self.merge and not conflicts: - self.stdout.write("No conflicts detected to merge.") - return - - # If they want to merge and there is something to merge, then - # divert into the merge code - if self.merge and conflicts: - return self.handle_merge(loader, conflicts) - - state = loader.project_state() - - # NOTE: customization. Passing graph to autodetector. - sql_graph = build_current_graph() - - # Set up autodetector - autodetector = MigrationAutodetector( - state, - ProjectState.from_apps(apps), - InteractiveMigrationQuestioner(specified_apps=app_labels, dry_run=self.dry_run), - sql_graph, - ) - - # If they want to make an empty migration, make one for each app - if self.empty: - if not app_labels: - raise CommandError("You must supply at least one app label when using --empty.") - # Make a fake changes() result we can pass to arrange_for_graph - changes = { - app: [Migration("custom", app)] - for app in app_labels - } - changes = autodetector.arrange_for_graph( - changes=changes, - graph=loader.graph, - migration_name=self.migration_name, - ) - self.write_migration_files(changes) - return - - # Detect changes - changes = autodetector.changes( - graph=loader.graph, - trim_to_apps=app_labels or None, - convert_apps=app_labels or None, - migration_name=self.migration_name, - ) - - if not changes: - # No changes? Tell them. - if self.verbosity >= 1: - if len(app_labels) == 1: - self.stdout.write("No changes detected in app '%s'" % app_labels.pop()) - elif len(app_labels) > 1: - self.stdout.write("No changes detected in apps '%s'" % ("', '".join(app_labels))) - else: - self.stdout.write("No changes detected") - - if self.exit_code: - sys.exit(1) - else: - self.write_migration_files(changes) - if check_changes: - sys.exit(1) From 0fef58f658997a45d1cce0f70df43a648830d0d1 Mon Sep 17 00:00:00 2001 From: Michal Byrtus Date: Thu, 5 Sep 2024 15:14:39 +0200 Subject: [PATCH 2/2] Django 5.1 support --- migrate_sql/autodetector.py | 14 +++++++------- 1 file changed, 7 insertions(+), 7 deletions(-) diff --git a/migrate_sql/autodetector.py b/migrate_sql/autodetector.py index cba24c0..e8a3a34 100644 --- a/migrate_sql/autodetector.py +++ b/migrate_sql/autodetector.py @@ -1,17 +1,18 @@ +from django.db.migrations.autodetector import OperationDependency from django.db.migrations.operations import RunSQL from django.utils.datastructures import OrderedSet +from migrate_sql.graph import SQLStateGraph, build_current_graph from migrate_sql.operations import ( AlterSQL, - ReverseAlterSQL, + AlterSQLState, CreateSQL, DeleteSQL, - AlterSQLState, + ReverseAlterSQL, ) -from migrate_sql.graph import SQLStateGraph, build_current_graph -class SQLBlob(object): +class SQLBlob: pass @@ -55,7 +56,7 @@ def is_sql_equal(sqls1, sqls2): if len(sqls1) != len(sqls2): return False - for sql1, sql2 in zip(sqls1, sqls2): + for sql1, sql2 in zip(sqls1, sqls2, strict=False): sql1, params1 = _sql_params(sql1) sql2, params2 = _sql_params(sql2) if sql1 != sql2 or params1 != params2: @@ -144,10 +145,9 @@ def add_sql_operation(self, app_label, sql_name, operation, dependencies): sequential operations. """ deps = [ - (dp[0], SQL_BLOB, dp[1], self._sql_operations.get(dp)) + OperationDependency(dp[0], SQL_BLOB, dp[1], self._sql_operations.get(dp)) for dp in dependencies ] - self.add_operation(app_label, operation, dependencies=deps) self._sql_operations[(app_label, sql_name)] = operation