Kaynağa Gözat

Closes #23314: Introduce NetBoxTestRunner to ensure --keepdb is always safe

Jeremy Stretch 2 gün önce
ebeveyn
işleme
8385e4f746

+ 2 - 2
.claude/skills/run-tests/SKILL.md

@@ -44,8 +44,8 @@ NETBOX_CONFIGURATION=netbox.configuration_testing python netbox/manage.py test d
 
 Speed options:
 
-- `--keepdb` — skip DB rebuild between runs (safe for most iterative work)
-- `--parallel` — run tests in parallel across CPU cores (used in CI; don't combine with `--keepdb` without testing first)
+- `--keepdb` — reuse the test DB between runs; it's rebuilt automatically when migrations change
+- `--parallel` — run tests in parallel across CPU cores (used in CI together with `--keepdb`)
 - `--failfast` — stop on first failure
 - `-v 2` — print each test name as it runs
 

+ 12 - 18
.github/workflows/ci.yml

@@ -157,26 +157,15 @@ jobs:
       - name: Collect static files
         run: python netbox/manage.py collectstatic --no-input
 
-      - name: Fingerprint the schema inputs
+      # The fingerprint NetBoxTestRunner stamps on kept test databases (the migrations of all installed apps).
+      - name: Fingerprint the migrations
         id: schema
         if: env.RELEASE_PR != 'true'
         run: |
-          # Path and blob of every file that shapes the migrated schema, so renames count too.
-          inputs=$(git ls-files --stage -- \
-            .github/workflows/ci.yml \
-            requirements.txt \
-            netbox/netbox/settings.py \
-            netbox/netbox/configuration_testing.py \
-            netbox/core/signals.py \
-            ':(glob)netbox/**/migrations/**' \
-            ':(glob)netbox/*/fields.py' \
-            netbox/netbox/models/ltree.py \
-            netbox/ipam/lookups.py \
-            netbox/utilities/query_functions.py \
-            netbox/utilities/ltree.py \
-            netbox/utilities/migration.py \
-            netbox/utilities/mptt_to_ltree.py)
-          echo "hash=$(sha256sum <<< "$inputs" | cut -c1-32)" >> "$GITHUB_OUTPUT"
+          hash=$(python netbox/manage.py shell --no-imports \
+            -c "from utilities.testing.runner import get_schema_fingerprint; print(get_schema_fingerprint())")
+          [[ "$hash" =~ ^[0-9a-f]{64}$ ]] || { echo "::error::Unexpected fingerprint output: $hash"; exit 1; }
+          echo "hash=$hash" >> "$GITHUB_OUTPUT"
 
       - name: Restore the migrated test database
         id: test-db
@@ -184,7 +173,8 @@ jobs:
         uses: actions/cache/restore@55cc8345863c7cc4c66a329aec7e433d2d1c52a9  # v6.1.0
         with:
           path: ${{ runner.temp }}/test_netbox.dump
-          key: test-db-py${{ matrix.python-version }}-${{ steps.schema.outputs.hash }}
+          # Bump the version when the build procedure changes; keep pg in sync with the postgres service image.
+          key: test-db-v1-pg18-py${{ matrix.python-version }}-${{ steps.schema.outputs.hash }}
           # Push runs migrate from scratch, so every merge re-validates the migration chain.
           lookup-only: ${{ github.event_name == 'push' }}
 
@@ -206,6 +196,10 @@ jobs:
           docker exec "$POSTGRES_CONTAINER" createdb -U netbox test_netbox
           docker exec -i "$POSTGRES_CONTAINER" pg_restore -U netbox -d test_netbox --exit-on-error \
             < "$RUNNER_TEMP/test_netbox.dump"
+          # Mark the restored database as current; NetBoxTestRunner otherwise rebuilds it. The cache key
+          # includes this same fingerprint, so a restored dump always matches the current migrations.
+          python netbox/manage.py shell --no-imports \
+            -c "from utilities.testing.runner import stamp_test_database; stamp_test_database()"
 
       - name: Cache the migrated test database
         if: env.RELEASE_PR != 'true' && steps.test-db.outputs.cache-hit != 'true'

+ 3 - 1
docs/development/getting-started.md

@@ -163,12 +163,14 @@ cd netbox/
 python manage.py test
 ```
 
-In cases where you haven't made any changes to the database schema (which is typical), you can append the `--keepdb` argument to this command to reuse the test database between runs. This cuts down on the time it takes to run the test suite since the database doesn't have to be rebuilt each time. (Note that this argument will cause errors if you've modified any model fields since the previous test run.)
+Building the test database from scratch can take several minutes, so it's recommended to always append the `--keepdb` argument to reuse the test database between runs:
 
 ```no-highlight
 python manage.py test --keepdb
 ```
 
+NetBox's test runner records a fingerprint of all migrations on the test database. If any migration changes (e.g. after switching branches or editing a migration), the test database is rebuilt automatically on the next run. If you encounter errors that don't seem related to your changes (for example, after interrupting a test run), run the test suite once without `--keepdb` to force a rebuild.
+
 You can also reduce testing time by enabling parallel test execution with the `--parallel` flag. (By default, this will run as many parallel tests as you have processors. To avoid sluggishness, it's a good idea to specify a lower number of parallel tests.) This flag can be combined with `--keepdb`, although if you encounter any strange errors, try running the test suite again with parallelization disabled.
 
 ```no-highlight

+ 2 - 0
netbox/netbox/settings.py

@@ -648,6 +648,8 @@ MESSAGE_TAGS = {
 
 DEFAULT_AUTO_FIELD = 'django.db.models.BigAutoField'
 
+TEST_RUNNER = 'utilities.testing.runner.NetBoxTestRunner'
+
 SERIALIZATION_MODULES = {
     'json': 'utilities.serializers.json',
 }

+ 128 - 0
netbox/utilities/testing/runner.py

@@ -0,0 +1,128 @@
+import hashlib
+import importlib.util
+import re
+import sys
+from pathlib import Path
+
+from django.apps import apps
+from django.db import DEFAULT_DB_ALIAS, DatabaseError, connections
+from django.db.migrations.loader import MigrationLoader
+from django.test.runner import DiscoverRunner
+
+__all__ = (
+    'NetBoxTestRunner',
+    'get_migration_files',
+    'get_schema_fingerprint',
+    'stamp_test_database',
+)
+
+FINGERPRINT_PREFIX = 'netbox-schema:'
+
+# Suffixes of editor backup and patch leftover files to exclude from the fingerprint
+IGNORED_SUFFIXES = ('~', '.bak', '.orig', '.rej')
+
+
+def get_migration_files():
+    """
+    Return a mapping of names to paths for all files (including data files) within the migrations package of each
+    installed app, including plugins. Hidden files (e.g. editor swap files), __pycache__, and backup files are ignored.
+    """
+    files = {}
+    for app_config in apps.get_app_configs():
+        module_name, _ = MigrationLoader.migrations_module(app_config.label)
+        if module_name and (spec := importlib.util.find_spec(module_name)) and spec.submodule_search_locations:
+            for location in spec.submodule_search_locations:
+                for path in Path(location).rglob('*'):
+                    relative_path = path.relative_to(location)
+                    if (
+                        path.is_file()
+                        and '__pycache__' not in relative_path.parts
+                        and not any(part.startswith('.') for part in relative_path.parts)
+                        and not path.name.endswith(IGNORED_SUFFIXES)
+                    ):
+                        files[f'{module_name}/{relative_path.as_posix()}'] = path
+    return files
+
+
+def get_schema_fingerprint():
+    """
+    Return a hash of the names and contents of all files returned by get_migration_files().
+
+    Migrations may invoke code elsewhere (e.g. custom fields or trigger SQL) which can change without any change to the
+    migrations themselves. Such changes are deliberately not tracked: they must be accompanied by a new migration to
+    take effect on existing databases, and until then a reused test database reflects what an upgrade would produce.
+    """
+    files = get_migration_files()
+    digest = hashlib.sha256()
+    for name in sorted(files):
+        digest.update(name.encode() + b'\0' + files[name].read_bytes() + b'\0')
+    return digest.hexdigest()
+
+
+def _write_stamp(connection, db_name, stamp):
+    with connection.creation._nodb_cursor() as cursor:
+        cursor.execute(f"COMMENT ON DATABASE {connection.ops.quote_name(db_name)} IS '{stamp}'")
+
+
+def stamp_test_database(alias=DEFAULT_DB_ALIAS):
+    """
+    Mark an existing test database as matching the current migrations, so that NetBoxTestRunner will reuse it with
+    --keepdb. Used by CI after restoring a cached test database.
+    """
+    connection = connections[alias]
+    _write_stamp(connection, connection.creation._get_test_db_name(), FINGERPRINT_PREFIX + get_schema_fingerprint())
+
+
+class NetBoxTestRunner(DiscoverRunner):
+    """
+    Extends Django's test runner to make --keepdb safe to use by default. Each kept test database is stamped with a
+    fingerprint of all migrations (see get_schema_fingerprint()). If the database's stamp is missing or does not
+    match the current fingerprint (e.g. after switching branches or editing a migration), the database and any clones
+    used by --parallel are rebuilt from scratch, rather than being migrated forward (or in the case of clones, reused
+    as-is).
+    """
+    def setup_databases(self, **kwargs):
+        if not self.keepdb:
+            return super().setup_databases(**kwargs)
+
+        stamp = FINGERPRINT_PREFIX + get_schema_fingerprint()
+        if (aliases := kwargs.get('aliases')) is None:
+            aliases = connections
+        aliases = [alias for alias in aliases if not connections[alias].settings_dict['TEST'].get('MIRROR')]
+        for alias in aliases:
+            self._discard_stale_test_databases(connections[alias], stamp)
+
+        old_config = super().setup_databases(**kwargs)
+
+        for alias in aliases:
+            connection = connections[alias]
+            _write_stamp(connection, connection.settings_dict['NAME'], stamp)
+
+        return old_config
+
+    def _discard_stale_test_databases(self, connection, stamp):
+        creation = connection.creation
+        test_db_name = creation._get_test_db_name()
+        clone_pattern = re.compile(rf'{re.escape(test_db_name)}_\d+')
+
+        with creation._nodb_cursor() as cursor:
+            cursor.execute("SELECT datname, shobj_description(oid, 'pg_database') FROM pg_catalog.pg_database")
+            databases = dict(cursor.fetchall())
+        if databases.get(test_db_name) == stamp:
+            return
+
+        to_drop = [name for name in databases if name == test_db_name or clone_pattern.fullmatch(name)]
+        if to_drop and self.verbosity >= 1:
+            creation.log(
+                f"Test database for alias '{connection.alias}' does not match the current migrations; "
+                f"discarding {', '.join(sorted(to_drop))}..."
+            )
+        for name in to_drop:
+            try:
+                creation._destroy_test_db(name, self.verbosity)
+            except DatabaseError as e:
+                creation.log(
+                    f"Got an error discarding test database {name}: {e}\n"
+                    f"Close any other connections to it and try again."
+                )
+                sys.exit(2)

+ 204 - 0
netbox/utilities/tests/test_runner.py

@@ -0,0 +1,204 @@
+import tempfile
+from pathlib import Path
+from unittest import mock
+
+from django.db import OperationalError
+from django.test import SimpleTestCase
+from django.test.runner import DiscoverRunner
+
+from utilities.testing import runner
+from utilities.testing.runner import NetBoxTestRunner
+
+STAMP = f'{runner.FINGERPRINT_PREFIX}abc'
+
+
+class SchemaFingerprintTestCase(SimpleTestCase):
+
+    def test_migration_files(self):
+        files = runner.get_migration_files()
+
+        self.assertIn('dcim.migrations/0001_squashed.py', files)
+        # Plugin migrations
+        self.assertIn('netbox.tests.dummy_plugin.migrations/0001_initial.py', files)
+        # Data files loaded by migrations
+        self.assertIn('dcim.migrations/initial_data/module_type_profiles/cpu.json', files)
+        self.assertFalse([name for name in files if '__pycache__' in name])
+
+    def test_migration_files_ignore_hidden_and_backup_files(self):
+        module_name = 'netbox.tests.dummy_plugin.migrations'
+        find_spec = runner.importlib.util.find_spec
+
+        with tempfile.TemporaryDirectory() as tmpdir:
+            for name in (
+                '0001_initial.py',
+                'data/profile.json',
+                '.0001_initial.py.swp',
+                '.hidden/0002_foo.py',
+                '0001_initial.py~',
+                '0001_initial.py.orig',
+                '0001_initial.py.rej',
+                '0001_initial.py.bak',
+            ):
+                path = Path(tmpdir) / name
+                path.parent.mkdir(parents=True, exist_ok=True)
+                path.touch()
+
+            def mock_find_spec(name, *args, **kwargs):
+                if name == module_name:
+                    return mock.Mock(submodule_search_locations=[tmpdir])
+                return find_spec(name, *args, **kwargs)
+
+            with mock.patch.object(runner.importlib.util, 'find_spec', side_effect=mock_find_spec):
+                files = runner.get_migration_files()
+
+        plugin_files = sorted(name for name in files if name.startswith(f'{module_name}/'))
+        self.assertEqual(plugin_files, [f'{module_name}/0001_initial.py', f'{module_name}/data/profile.json'])
+
+    def test_fingerprint_reflects_names_and_contents(self):
+        with tempfile.TemporaryDirectory() as tmpdir:
+            path = Path(tmpdir) / 'file.py'
+            path.write_text('foo')
+
+            with mock.patch.object(runner, 'get_migration_files', return_value={'a': path}):
+                fingerprint = runner.get_schema_fingerprint()
+                self.assertEqual(runner.get_schema_fingerprint(), fingerprint)
+
+                # Changing a file's contents changes the fingerprint
+                path.write_text('bar')
+                self.assertNotEqual(runner.get_schema_fingerprint(), fingerprint)
+                fingerprint = runner.get_schema_fingerprint()
+
+            # Renaming a file changes the fingerprint
+            with mock.patch.object(runner, 'get_migration_files', return_value={'b': path}):
+                self.assertNotEqual(runner.get_schema_fingerprint(), fingerprint)
+
+
+def get_mock_connection(name='netbox', mirror=None, databases=None):
+    """
+    Return a mock database connection whose server reports the given {name: comment} databases.
+    """
+    connection = mock.MagicMock(alias='default', settings_dict={'NAME': name, 'TEST': {'MIRROR': mirror}})
+    connection.ops.quote_name.side_effect = lambda name: f'"{name}"'
+    # Like Django, derive the test database name from the connection's current database name
+    connection.creation._get_test_db_name.side_effect = lambda: f"test_{connection.settings_dict['NAME']}"
+    cursor = connection.creation._nodb_cursor.return_value.__enter__.return_value
+    cursor.fetchall.return_value = list((databases or {}).items())
+    return connection
+
+
+class DiscardStaleTestDatabasesTestCase(SimpleTestCase):
+
+    def get_dropped_databases(self, databases):
+        connection = get_mock_connection(databases=databases)
+        NetBoxTestRunner(keepdb=True, verbosity=0)._discard_stale_test_databases(connection, STAMP)
+        return {c.args[0] for c in connection.creation._destroy_test_db.call_args_list}
+
+    def test_matching_stamp(self):
+        dropped = self.get_dropped_databases({
+            'test_netbox': STAMP,
+            'test_netbox_1': None,
+        })
+        self.assertEqual(dropped, set())
+
+    def test_stale_stamp(self):
+        dropped = self.get_dropped_databases({
+            'netbox': None,
+            'test_netbox': f'{runner.FINGERPRINT_PREFIX}old',
+            'test_netbox_1': None,
+            'test_netbox_12': None,
+            'test_netbox_branching': None,
+            'test_netbox_1_old': None,
+        })
+        self.assertEqual(dropped, {'test_netbox', 'test_netbox_1', 'test_netbox_12'})
+
+    def test_missing_stamp(self):
+        dropped = self.get_dropped_databases({
+            'test_netbox': None,
+            'test_netbox_1': None,
+        })
+        self.assertEqual(dropped, {'test_netbox', 'test_netbox_1'})
+
+    def test_missing_database(self):
+        dropped = self.get_dropped_databases({
+            'test_netbox_1': None,
+        })
+        self.assertEqual(dropped, {'test_netbox_1'})
+
+    def test_drop_error(self):
+        connection = get_mock_connection(databases={'test_netbox': None})
+        connection.creation._destroy_test_db.side_effect = OperationalError('database is being accessed')
+
+        with self.assertRaises(SystemExit) as cm:
+            NetBoxTestRunner(keepdb=True, verbosity=0)._discard_stale_test_databases(connection, STAMP)
+        self.assertEqual(cm.exception.code, 2)
+        self.assertIn('Close any other connections', connection.creation.log.call_args.args[0])
+
+
+@mock.patch.object(runner, 'get_schema_fingerprint', return_value='abc')
+@mock.patch.object(runner, '_write_stamp')
+@mock.patch.object(NetBoxTestRunner, '_discard_stale_test_databases')
+class SetupDatabasesTestCase(SimpleTestCase):
+
+    def setUp(self):
+        self.connections = {
+            'default': get_mock_connection(),
+            'replica': get_mock_connection(mirror='default'),
+        }
+
+        def setup_databases(runner, **kwargs):
+            # Mimic Django pointing each set up connection at its test database
+            for alias in kwargs['aliases'] if kwargs['aliases'] is not None else self.connections:
+                settings_dict = self.connections[alias].settings_dict
+                settings_dict['NAME'] = f"test_{settings_dict['NAME']}"
+            return 'old_config'
+
+        patchers = (
+            mock.patch.object(runner, 'connections', self.connections),
+            mock.patch.object(DiscoverRunner, 'setup_databases', autospec=True, side_effect=setup_databases),
+        )
+        for patcher in patchers:
+            patcher.start()
+            self.addCleanup(patcher.stop)
+
+    def test_without_keepdb(self, mock_discard, mock_write_stamp, mock_fingerprint):
+        old_config = NetBoxTestRunner(keepdb=False, verbosity=0).setup_databases(aliases={'default'})
+
+        self.assertEqual(old_config, 'old_config')
+        mock_discard.assert_not_called()
+        mock_write_stamp.assert_not_called()
+
+    def test_with_keepdb(self, mock_discard, mock_write_stamp, mock_fingerprint):
+        default = self.connections['default']
+        old_config = NetBoxTestRunner(keepdb=True, verbosity=0).setup_databases(aliases={'default', 'replica'})
+
+        self.assertEqual(old_config, 'old_config')
+        # Mirrors are skipped
+        mock_discard.assert_called_once_with(default, STAMP)
+        # The stamp is written to the test database, not the original
+        mock_write_stamp.assert_called_once_with(default, 'test_netbox', STAMP)
+
+    def test_all_aliases(self, mock_discard, mock_write_stamp, mock_fingerprint):
+        NetBoxTestRunner(keepdb=True, verbosity=0).setup_databases(aliases=None)
+
+        mock_discard.assert_called_once_with(self.connections['default'], STAMP)
+        mock_write_stamp.assert_called_once_with(self.connections['default'], 'test_netbox', STAMP)
+
+    def test_no_aliases(self, mock_discard, mock_write_stamp, mock_fingerprint):
+        # A suite without any database tests (e.g. only SimpleTestCases) must not touch any database
+        NetBoxTestRunner(keepdb=True, verbosity=0).setup_databases(aliases=set())
+
+        mock_discard.assert_not_called()
+        mock_write_stamp.assert_not_called()
+
+
+class StampTestDatabaseTestCase(SimpleTestCase):
+
+    @mock.patch.object(runner, 'get_schema_fingerprint', return_value='abc')
+    def test_stamp_test_database(self, mock_fingerprint):
+        connection = get_mock_connection()
+
+        with mock.patch.object(runner, 'connections', {'default': connection}):
+            runner.stamp_test_database()
+
+        cursor = connection.creation._nodb_cursor.return_value.__enter__.return_value
+        cursor.execute.assert_called_once_with(f'COMMENT ON DATABASE "test_netbox" IS \'{STAMP}\'')