diff --git a/.devcontainer/scripts/setup.sh b/.devcontainer/scripts/setup.sh index f2af2ece..0b2087e8 100755 --- a/.devcontainer/scripts/setup.sh +++ b/.devcontainer/scripts/setup.sh @@ -144,9 +144,8 @@ if ! command -v gh >/dev/null 2>&1; then (type -p wget >/dev/null || apt-get install -y -qq wget) \ && install -d -m 755 /etc/apt/keyrings \ && out=$(mktemp) \ - && wget -qO "$out" https://cli.github.com/packages/githubcli-archive-keyring.gpg \ - && cat "$out" | tee /etc/apt/keyrings/githubcli-archive-keyring.gpg > /dev/null \ - && chmod go+r /etc/apt/keyrings/githubcli-archive-keyring.gpg \ + && wget -q --max-redirect=0 -O "$out" https://cli.github.com/packages/githubcli-archive-keyring.gpg \ + && install -m 644 "$out" /etc/apt/keyrings/githubcli-archive-keyring.gpg \ && echo "deb [arch=$(dpkg --print-architecture) signed-by=/etc/apt/keyrings/githubcli-archive-keyring.gpg] https://cli.github.com/packages stable main" | tee /etc/apt/sources.list.d/github-cli.list > /dev/null \ && apt-get update -qq \ && apt-get install -y -qq gh \ diff --git a/.github/copilot-instructions.md b/.github/copilot-instructions.md index c607099d..15145c72 100644 --- a/.github/copilot-instructions.md +++ b/.github/copilot-instructions.md @@ -11,8 +11,10 @@ Requires NetBox ≥ 4.3.0 and Python ≥ 3.12. Licensed under Apache-2.0 (REUSE- This follows the standard [NetBox plugin pattern](https://netboxlabs.com/docs/netbox/en/stable/plugins/development/): - **`models.py`**: Defines `InterfaceNameRule`. A rule can select an exact module type or a regex pattern, add parent, device, and platform scopes, and describe flat or channelized breakout output. -- **`signals.py`**: Connects `pre_save` and `post_save` for `dcim.Module` and `dcim.Device` and passes each save to `rename_triggers.py`. The receivers hold no state and make no decision. It intentionally does not connect to `dcim.Interface` because NetBox creates module interfaces with `bulk_create()`. It also connects the optional LibreNMS prediction signal when that plugin is installed. -- **`rename_triggers.py`**: Owns the rename-trigger lifecycle. It reads the previous state before a save and lets a read error fail the save. After the save it decides whether the save is a rename trigger and schedules one reapply per module or device per transaction with `transaction.on_commit()`. The reapply compares the earliest previous state with the committed row, and it catches and logs failures at that boundary. +- **`signals.py`**: Connects `pre_save` and `post_save` for `dcim.Module`, `dcim.ModuleBay` and `dcim.Device` and passes each save to `rename_triggers.py` with the alias of the save, which must be the write alias. The receivers hold no state and make no decision. It intentionally does not connect to `dcim.Interface` because NetBox creates module interfaces with `bulk_create()`. It also connects the optional LibreNMS prediction signal when that plugin is installed. +- **`rename_triggers.py`**: Owns the rename-trigger lifecycle. It reads the previous state before a save and lets a read error fail the save. After the save it decides whether the save is a rename trigger and schedules one reapply per module or device per transaction with `transaction.on_commit()` on the connection of the save, so the plan runs after that connection commits. The reapply compares the earliest previous state with the committed row, and it catches and logs failures at that boundary. +- **`transactions.py`**: The one owner of database connections and transaction state. Each plugin write runs in a write scope on `default`, and in a netbox-branching branch also on the branch connection. +- **`branching.py`**: The one module that imports netbox-branching. It checks its version at startup, gives a job the identity of its branch, and marks each merge, revert and sync so that the rename triggers do nothing while it replays changes. - **`rule_selection.py`**: Loads and fingerprints enabled rules, separates exact and regex candidates, applies scope priority, and pins one cached snapshot across batch work. - **`name_template.py`**: Owns the name-template language, its template-variable catalogue, and evaluation. - **`naming.py`**: Builds template-variable values from the module-bay hierarchy. @@ -78,7 +80,7 @@ netbox-test netbox_interface_name_rules/tests/test_views.py::TestClassName::test TEST_DB_NAME=test_netbox_interface_name_rules TEST_REDIS_HOST=redis pytest netbox_interface_name_rules ``` -`pyproject.toml` adds `-n auto` and coverage options. Do not pass your own `-n`. +`pyproject.toml` adds `-n auto` and coverage options. Do not pass your own `-n`. A local run prints coverage but does not fail on it: the 97% gate runs in CI on the combined data of two legs. ## REUSE/SPDX compliance diff --git a/.github/workflows/coverage-badge.yaml b/.github/workflows/coverage-badge.yaml index 6e1df179..b4eeee2a 100644 --- a/.github/workflows/coverage-badge.yaml +++ b/.github/workflows/coverage-badge.yaml @@ -27,7 +27,7 @@ jobs: steps: - name: Download coverage report - uses: actions/download-artifact@3e5f45b2cfb9172054b4087a40e8e0b5a5461e7c # v4 + uses: actions/download-artifact@3e5f45b2cfb9172054b4087a40e8e0b5a5461e7c # v8.0.1 with: name: coverage-report run-id: ${{ github.event.workflow_run.id }} diff --git a/.github/workflows/mkdocs.yaml b/.github/workflows/mkdocs.yaml index 46ecab36..1663f68b 100644 --- a/.github/workflows/mkdocs.yaml +++ b/.github/workflows/mkdocs.yaml @@ -41,7 +41,7 @@ jobs: python-version: '3.12' - name: Install dependencies - run: uv pip install --system --group docs + run: uv sync --locked --no-build --only-group docs - name: Save coverage report from gh-pages run: | @@ -59,7 +59,7 @@ jobs: fi - name: Deploy documentation - run: mkdocs gh-deploy --force --strict + run: uv run --no-sync --no-build mkdocs gh-deploy --force --strict - name: Restore coverage report if: always() diff --git a/.github/workflows/publish-pypi.yaml b/.github/workflows/publish-pypi.yaml index 2dd11b3f..bec74a24 100644 --- a/.github/workflows/publish-pypi.yaml +++ b/.github/workflows/publish-pypi.yaml @@ -35,9 +35,9 @@ jobs: with: python-version: "3.x" - name: Install pypa/build - run: uv pip install --system --group packaging + run: uv sync --locked --no-build --only-group packaging - name: Build a binary wheel and a source tarball - run: python3 -m build + run: uv run --no-sync --no-build python -m build - name: Store the distribution packages uses: actions/upload-artifact@043fb46d1a93c77aae656e7c1c64a875d1fc6a0a # v7.0.1 with: diff --git a/.github/workflows/release.yaml b/.github/workflows/release.yaml index 8414b309..bd474ac3 100644 --- a/.github/workflows/release.yaml +++ b/.github/workflows/release.yaml @@ -41,9 +41,9 @@ jobs: uses: astral-sh/setup-uv@c18668ad3cf93ea998bef934396af7bb5c839dc7 # v10.2.0 - name: Install dependencies - run: uv sync --group dev + run: uv sync --locked --no-build --only-group release --only-group packaging - name: Run semantic-release env: GH_TOKEN: ${{ secrets.RELEASE_TOKEN }} - run: uv run semantic-release version --changelog --push --vcs-release + run: uv run --no-sync --no-build semantic-release version --changelog --push --vcs-release diff --git a/.github/workflows/test.yaml b/.github/workflows/test.yaml index 95f08355..3e6bcba0 100644 --- a/.github/workflows/test.yaml +++ b/.github/workflows/test.yaml @@ -34,15 +34,16 @@ jobs: netbox-version: "v4.3.7" - python-version: "3.13" netbox-version: "v4.5.3" - # The one cell that reports coverage. Every coverage step below reads this flag, so the - # release that owns the report is named here and nowhere else. - coverage: true + # A coverage leg: the value names its data artifact, which the combine job downloads. + coverage: "netbox-v4.5.3" - python-version: "3.14" netbox-version: "v4.7.0" # The same release with netbox-branching, so the branch tests run against a real provisioned branch. - python-version: "3.14" netbox-version: "v4.7.0" netbox-branching: "1.2.1" + # The second coverage leg: only this cell runs the netbox-branching code paths. + coverage: "netbox-v4.7.0-branching" # Early warning for upcoming NetBox changes. Non-blocking: unreleased NetBox breaks # for reasons that are not ours, so this must not gate a PR. - python-version: "3.13" @@ -164,13 +165,72 @@ jobs: # instead of baseline drift. Bumping the top version means editing the matrix above, this # condition, and tests/query_counts.json together. The branch leg records too: netbox-branching adds queries. UPDATE_QUERY_COUNTS: ${{ (matrix.netbox-version != 'v4.7.0' || matrix.netbox-branching) && '1' || '' }} - # `addopts` turns coverage on for every invocation, but only one cell reports it. + # `addopts` collects coverage without the gate; the combine job enforces `fail_under`. COVERAGE_OPTION: ${{ !matrix.coverage && '--no-cov' || '' }} run: | pytest -n auto netbox_interface_name_rules -o pythonpath=../netbox/netbox $COVERAGE_OPTION - - name: Generate coverage report + - name: Upload coverage data if: matrix.coverage + uses: actions/upload-artifact@043fb46d1a93c77aae656e7c1c64a875d1fc6a0a # v7.0.1 + with: + name: coverage-data-${{ matrix.coverage }} + path: netbox-InterfaceNameRules-plugin/.coverage + include-hidden-files: true + if-no-files-found: error + retention-days: 1 + + # The 97% gate applies to the union of the coverage legs, because some code runs only with netbox-branching. + coverage: + needs: test-netbox + runs-on: ubuntu-latest + permissions: + contents: read + + steps: + # The same path as in the test job, so the absolute paths in the data files resolve here. + - name: Checkout code + uses: actions/checkout@3d3c42e5aac5ba805825da76410c181273ba90b1 # v7.0.1 + with: + path: netbox-InterfaceNameRules-plugin + persist-credentials: false + + - name: Install uv + uses: astral-sh/setup-uv@bec219d24cd3e171d82865faccec33120bb574f4 # v10.1.0 + + - name: Set up Python + uses: actions/setup-python@5fda3b95a4ea91299a34e894583c3862153e4b97 # v7.0.0 + with: + python-version: "3.14" + + # One download per leg by exact name, so a missing leg fails this step. + - name: Download coverage data of NetBox v4.5.3 + uses: actions/download-artifact@3e5f45b2cfb9172054b4087a40e8e0b5a5461e7c # v8.0.1 + with: + name: coverage-data-netbox-v4.5.3 + path: coverage-data/netbox-v4.5.3 + + - name: Download coverage data of NetBox v4.7.0 with netbox-branching + uses: actions/download-artifact@3e5f45b2cfb9172054b4087a40e8e0b5a5461e7c # v8.0.1 + with: + name: coverage-data-netbox-v4.7.0-branching + path: coverage-data/netbox-v4.7.0-branching + + - name: Install coverage + working-directory: netbox-InterfaceNameRules-plugin + run: | + uv pip install --system --only-binary=:all: --group ci-tests + + # `coverage report` enforces `fail_under` from pyproject.toml on the combined data. + - name: Combine and check coverage + working-directory: netbox-InterfaceNameRules-plugin + env: + COVERAGE_RCFILE: pyproject.toml + run: | + coverage combine ../coverage-data/netbox-v4.5.3/.coverage ../coverage-data/netbox-v4.7.0-branching/.coverage + coverage report + + - name: Generate coverage report working-directory: netbox-InterfaceNameRules-plugin env: COVERAGE_RCFILE: pyproject.toml @@ -181,7 +241,6 @@ jobs: coverage xml -o ../coverage-report/coverage.xml - name: Upload coverage report - if: matrix.coverage uses: actions/upload-artifact@043fb46d1a93c77aae656e7c1c64a875d1fc6a0a # v7.0.1 with: name: coverage-report @@ -189,7 +248,6 @@ jobs: retention-days: 7 - name: Upload coverage to Codecov - if: matrix.coverage uses: codecov/codecov-action@303a32d7a59b442fa8d48b6a1cc6825c09c847a5 # v7.1.1 with: files: coverage-report/coverage.xml diff --git a/CONTEXT.md b/CONTEXT.md index f3a69378..f600d79e 100644 --- a/CONTEXT.md +++ b/CONTEXT.md @@ -94,7 +94,7 @@ The complete path from a NetBox model save, through the committed callback, to t _Avoid_: Signal handler performance **Rename trigger**: -A saved change in NetBox after which the names a rule gives may be wrong, so the plugin must reapply its rules. The triggers are: a module is installed, a module's type changes, a module moves to another bay or device, an occupied module bay's position or name changes, and a device's virtual chassis or virtual-chassis position changes (the device joins a virtual chassis, leaves it, or gets a different position). A bay's name change is a trigger only when a template variable reads the name: a bay whose position is a template token takes its position from the trailing digits of its name. An edit of an empty bay is not a trigger. A module's type change also reaches each module nested in it whose rule changes, because a rule can be scoped to a parent module type. The triggers of one transaction cause one reapply plan, which reapplies each module and each device at most once. A device reapply also reapplies every module of its device that the plan did not already reapply for a module trigger. +A saved change in NetBox after which the names a rule gives may be wrong, so the plugin must reapply its rules. The triggers are: a module is installed, a module's type changes, a module moves to another bay or device, an occupied module bay's position or name changes, and a device's virtual chassis or virtual-chassis position changes (the device joins a virtual chassis, leaves it, or gets a different position). A bay's name change is a trigger only when a template variable reads the name: a bay whose position is a template token takes its position from the trailing digits of its name. An edit of an empty bay is not a trigger. A module's type change also reaches each module nested in it whose rule changes, because a rule can be scoped to a parent module type. A change that netbox-branching replays in a merge, a revert or a sync is not a rename trigger: the replayed changes already hold the names. The triggers of one transaction cause one reapply plan, which reapplies each module and each device at most once. A device reapply also reapplies every module of its device that the plan did not already reapply for a module trigger. _Avoid_: Signal, event **Reapply**: diff --git a/CONTRIBUTING.md b/CONTRIBUTING.md index 35054af4..5f73b764 100644 --- a/CONTRIBUTING.md +++ b/CONTRIBUTING.md @@ -55,6 +55,9 @@ with `test_`, and set `TEST_REDIS_HOST` to the Redis server that the tests can u TEST_DB_NAME=test_netbox_interface_name_rules TEST_REDIS_HOST=localhost pytest netbox_interface_name_rules ``` +A local run prints coverage but does not fail on it. The 97% gate runs in CI on the combined +coverage data of the NetBox v4.5.3 leg and the netbox-branching leg. + ### Commits We use [Conventional Commits](https://www.conventionalcommits.org/): diff --git a/README.md b/README.md index d171bf35..32557134 100644 --- a/README.md +++ b/README.md @@ -29,6 +29,7 @@ automatically apply renaming rules based on configurable templates. - **Breakout support**: create multiple channel interfaces from a single port (e.g., QSFP+ 4x10G) - **Scoping**: rules can be scoped to specific device types, parent module types, or be universal - **Bulk import/export**: YAML-based rule management via the UI or API +- **netbox-branching**: renames run in the active branch, and a merge, revert or sync keeps the names that it replays (netbox-branching 1.2.x on NetBox 4.7). A channel that the plugin kept at its old name is the exception: see [Limits in a branch](https://marcinpsk.github.io/netbox-InterfaceNameRules-plugin/configuration/#limits-in-a-branch) ## Supported scenarios diff --git a/docs/configuration.md b/docs/configuration.md index c089a1b1..6dc0754a 100644 --- a/docs/configuration.md +++ b/docs/configuration.md @@ -286,13 +286,103 @@ refuses another member, and sends an event for each one that it kept. **Apply Rules** is designed for **retroactive renames**. Interfaces installed after a matching rule is active are renamed automatically at install time. The web UI, the REST API and bulk import install modules inside a transaction. -A script or shell that creates a module outside a transaction gets no rename: -run Apply Rules after it, or wrap the install in `transaction.atomic()`. +A script or shell that creates a module outside a transaction on the interface +write connection gets no rename: run Apply Rules after it, or wrap the install in +`transaction.atomic(using=router.db_for_write(Interface))`. In a netbox-branching +branch, `transaction.atomic()` alone opens a transaction on `default`, not on the +branch connection, so the rename runs before NetBox creates the interfaces. The **Applicable** column shows ✓ only when at least one currently-installed interface **would actually change name** if the rule were applied. Rules where all matching interfaces are already correctly named show `—`. +## netbox-branching + +The plugin supports [netbox-branching](https://github.com/netboxlabs/netbox-branching) +1.2.x on NetBox 4.7. When netbox-branching is installed, NetBox does not start +with another release of it, and the error names the installed version. Without +netbox-branching, the plugin supports NetBox 4.3 to 4.7 and works as this guide +describes. + +### In a branch + +While a branch is active, the plugin reads and writes in that branch only: + +- A rename trigger in the branch renames the interfaces in the branch, after + NetBox commits the change. Each rename has a change record in the branch, so a + merge applies it to main. +- **Apply Rules** and the flat-to-channelized conversion change the interfaces + of the branch. A rule that exists only in the branch applies only there. +- **Run as Background Job** and **Convert as Background Job** run in the branch + that was active when you started the job. When that branch is not ready when + the job starts, for example because it was merged, the job fails and changes + nothing. +- A script or the shell must install a module in a transaction on the interface + write connection, as [Apply Rules and the Applicable Column](#apply-rules-and-the-applicable-column) + describes. + +In a branch, each plugin operation sets the PostgreSQL `lock_timeout` to 10 +seconds on the connection of the branch and on the connection of main. When the +operation ends, the plugin sets the earlier values again. A request in a branch +holds two PostgreSQL sessions, and PostgreSQL does not find a lock cycle through +the two sessions of one request. An operation that waits longer for a lock stops +with an error. A script can call an engine function, such as +`apply_device_interface_rules`, inside a transaction that the script holds. The +commit callbacks of the plugin then run when that transaction commits, with the +earlier `lock_timeout` of the session. Set `lock_timeout` for these callbacks in +your script. + +### Merge, revert and sync + +netbox-branching merges, reverts and syncs a branch: it replays the changes that +NetBox logged. The replayed changes already hold the interface names, so the +rename triggers do nothing while netbox-branching replays them. + +- A merge gives main the interface names of the branch, except a kept channel + (see [Limits in a branch](#limits-in-a-branch)). On main, it writes only the + replayed changes. +- A revert of the merge gives main the names from before the merge. +- A sync gives the branch the names of main, except a kept channel, and writes no + rename without a change record. A rule that exists only in the branch does not rename the interfaces + that the sync brought. Run **Apply Rules** in the branch after the sync to + apply it. + +After a merge, a revert or a sync that fails or stops early, for example a dry +run or a merge of a branch without changes, the next change is a rename trigger +again. + +### Limits in a branch + +- **A replay can rename a kept channel.** When the name that a rule gives a + channel subinterface is in use, the plugin keeps the old name of the channel + and renames its parent. NetBox renames the channels of a renamed parent when + the change commits, and the plugin then gives the kept channel its old name + again. A merge, a revert or a sync replays the rename of the parent, so + NetBox renames the channels again when the replay commits, and the plugin does + not act. After a merge, the kept channel on main then has the name from NetBox, + for example `et-0/0/1:2`, while the channel in the branch keeps `1:2`. After a + sync, the kept channel in the branch has the name from NetBox. A revert of the + merge gives the channel its name from before the merge. Rename such a channel + by hand when you want the name from the other side. +- **A background REST request runs on main.** NetBox runs a bulk REST request + with `background=true` as a background job, and that job does not keep the + active branch. The REST API of the rules refuses such a request while a branch + is active, before it writes. The other NetBox endpoints run it on main. For + example, modules that you install with such a request are installed on main, + and the plugin renames their interfaces on main. +- **A failed branch activation runs the request on main.** When NetBox cannot + activate the branch of a request, it continues the request on main. This + applies to all changes of the request, not only to the plugin. +- **No atomicity across the two connections.** netbox-branching records each + change of a branch on the connection of main. The plugin commits the connection + of the branch first and the connection of main second, as NetBox scripts and + netbox-branching do. When the second commit fails, or a NetBox callback fails + after the first commit, the branch keeps the renames, but the list of branch + changes in netbox-branching does not show them. A merge still applies them. +- **Connection pooling in transaction mode is not supported**, for example + PgBouncer with `pool_mode = transaction`. A session setting such as + `lock_timeout` does not stay with the session of such a pool. + ## Bulk Import Export existing rules or import new ones via **Interface Name Rules → Import**. diff --git a/docs/design/netbox-branching.md b/docs/design/netbox-branching.md index 972a1ea9..5d714a6a 100644 --- a/docs/design/netbox-branching.md +++ b/docs/design/netbox-branching.md @@ -174,13 +174,16 @@ Changes from r1 (round 1), carried into r2: on the branch connection (INR transactions.py:29-43). - Invariant: INR writes only through an alias of the open scope. -**`branching.py`** is the one module that imports `netbox_branching`, loaded only when it is -installed. `AppConfig.ready()` calls it. +**`branching.py`** is the one module that imports `netbox_branching`; an AST test enforces it. The +plugin imports `branching.py` on every install, and `branching.py` imports `netbox_branching` only +inside a function that runs when netbox-branching is installed. `AppConfig.ready()` calls its +version check. - Version gate: raise `ImproperlyConfigured` unless the installed netbox-branching is 1.2.x. - Replay suppression: wrap `Branch.merge`, `Branch.revert` and `Branch.sync` once (idempotent, `functools.wraps`). Each wrapper sets a ContextVar token and resets it in `finally`. - `replay_in_progress() -> bool`. + `replay_in_progress() -> bool`. `InterfaceNameRule.save()` skips its write-alias check while it is + true: a merge or revert started in an active branch replays a rule update on `default`. - Job identity: `branch_identity() -> str | None` (the active branch's schema id) and `activate_on(request, identity)`, which sets BR's branch cookie on a synthetic request. diff --git a/docs/examples.md b/docs/examples.md index 46c15540..45a620aa 100644 --- a/docs/examples.md +++ b/docs/examples.md @@ -410,6 +410,10 @@ Arista modular/multi-chassis naming uses `Ethernet{slot}/{port}`. The device typ for dev in Device.objects.filter(virtual_chassis__isnull=False): apply_device_interface_rules(dev) ``` + In a netbox-branching branch, the function sets a `lock_timeout` of 10 seconds + while it runs. When you call it inside a transaction that you hold, its commit + callbacks run when your transaction commits, with the earlier `lock_timeout` of + the session. Set `lock_timeout` for these callbacks in your script. --- diff --git a/docs/index.md b/docs/index.md index 1597897f..ae4826fd 100644 --- a/docs/index.md +++ b/docs/index.md @@ -23,6 +23,7 @@ automatically apply renaming rules based on configurable templates. - **Scoping**: rules can target specific device types, parent module types, platforms, or be universal - **Build Rule tester**: preview module and device-interface names before saving. Module rules also preview matching installed interfaces. - **Apply Rules**: batch rename existing interfaces with live preview and background job support +- **netbox-branching**: renames run in the active branch, and a merge, revert or sync keeps the names that it replays (netbox-branching 1.2.x on NetBox 4.7). A channel that the plugin kept at its old name is the exception: see [Limits in a branch](configuration.md#limits-in-a-branch) ## Supported Scenarios diff --git a/docs/installation.md b/docs/installation.md index 2bffee74..687f95cf 100644 --- a/docs/installation.md +++ b/docs/installation.md @@ -4,6 +4,8 @@ - NetBox ≥ 4.3.0 - Python ≥ 3.12 +- Optional: netbox-branching 1.2.x, on NetBox 4.7. See + [netbox-branching](configuration.md#netbox-branching). ## Install from PyPI @@ -19,6 +21,20 @@ Add to your NetBox `configuration.py`: PLUGINS = ["netbox_interface_name_rules"] ``` +## Before You Upgrade + +Let the queued **Run as Background Job** and **Convert as Background Job** jobs +finish before you upgrade the plugin. A release can change the data that a job +stores, and a job that an earlier release queued then fails after the upgrade. +The release that adds netbox-branching support stores the branch of each job. +Start a failed job again after the upgrade. + +With netbox-branching, a script can call an engine function, such as +`apply_device_interface_rules`, inside a transaction that the script holds. The +commit callbacks of the plugin then run when that transaction commits, with the +earlier `lock_timeout` of the session. Set `lock_timeout` for these callbacks in +your script. See [netbox-branching](configuration.md#netbox-branching). + ## Run Database Migrations The migration audits every existing nonempty **Module Type Pattern** used by a diff --git a/netbox_interface_name_rules/__init__.py b/netbox_interface_name_rules/__init__.py index e2875217..ff29869a 100644 --- a/netbox_interface_name_rules/__init__.py +++ b/netbox_interface_name_rules/__init__.py @@ -1,6 +1,5 @@ # SPDX-License-Identifier: Apache-2.0 # Copyright (C) 2025 Marcin Zieba -from django.apps import apps from netbox.plugins import PluginConfig __version__ = "1.7.0" @@ -23,14 +22,11 @@ class InterfaceNameRulesConfig(PluginConfig): author_email = "marcinpsk@gmail.com" def ready(self): - """Connect signal handlers after all apps are loaded, and check netbox-branching when it is installed.""" + """Connect signal handlers after all apps are loaded, and prepare for netbox-branching when it is installed.""" super().ready() - from . import signals # registers the post_save handler + from . import branching, signals # signals registers the post_save handler - if apps.is_installed("netbox_branching"): - from . import branching - - branching.check_installed_version() + branching.ready() config = InterfaceNameRulesConfig diff --git a/netbox_interface_name_rules/api/views.py b/netbox_interface_name_rules/api/views.py index 565ed5bd..bb156b73 100644 --- a/netbox_interface_name_rules/api/views.py +++ b/netbox_interface_name_rules/api/views.py @@ -1,14 +1,26 @@ # SPDX-License-Identifier: Apache-2.0 # Copyright (C) 2025 Marcin Zieba from netbox.api.viewsets import NetBoxModelViewSet +from rest_framework.exceptions import ValidationError +from netbox_interface_name_rules import branching from netbox_interface_name_rules.models import InterfaceNameRule from .serializers import InterfaceNameRuleSerializer +BACKGROUND_IN_A_BRANCH = ( + "A background request runs on main, not in the active branch. Send the request without background=true." +) + class InterfaceNameRuleViewSet(NetBoxModelViewSet): """REST API viewset for InterfaceNameRule.""" queryset = InterfaceNameRule.objects.all() serializer_class = InterfaceNameRuleSerializer + + def _enqueue_bulk_job(self, request, *args, **kwargs): + """Refuse a background request in a branch before NetBox enqueues it: its job runs on main.""" + if branching.branch_identity() is not None: + raise ValidationError(BACKGROUND_IN_A_BRANCH) + return super()._enqueue_bulk_job(request, *args, **kwargs) diff --git a/netbox_interface_name_rules/branching.py b/netbox_interface_name_rules/branching.py index 937e74bd..dcf7ab48 100644 --- a/netbox_interface_name_rules/branching.py +++ b/netbox_interface_name_rules/branching.py @@ -1,13 +1,23 @@ # SPDX-License-Identifier: Apache-2.0 # Copyright (C) 2025 Marcin Zieba -"""The one module that uses netbox-branching. The plugin loads it only when netbox-branching is installed.""" +"""The one module that uses netbox-branching. It imports netbox-branching only in a function that needs it.""" + +import functools +from contextvars import ContextVar from django.apps import apps from django.core.exceptions import ImproperlyConfigured from packaging.version import Version +APP_LABEL = "netbox_branching" SUPPORTED_RELEASE = (1, 2) SUPPORTED_SERIES = ".".join(map(str, SUPPORTED_RELEASE)) + ".x" +# The methods of netbox-branching's Branch that replay logged changes. +REPLAYING_METHODS = ("merge", "revert", "sync") +# The attribute that marks a method this module wrapped. +REPLAY_MARK = "marks_a_replay" + +_replaying = ContextVar("netbox_interface_name_rules_replay", default=False) def check_version(version: str) -> None: @@ -18,6 +28,58 @@ def check_version(version: str) -> None: ) -def check_installed_version() -> None: - """Check the installed netbox-branching when NetBox starts.""" - check_version(apps.get_app_config("netbox_branching").version) +def installed_version() -> str: + """Return the version of the installed netbox-branching.""" + return apps.get_app_config(APP_LABEL).version + + +def ready() -> None: + """When netbox-branching is installed, check its version and mark each of its replays in the context that runs it.""" + if not apps.is_installed(APP_LABEL): + return + check_version(installed_version()) + from netbox_branching.models import Branch + + for name in REPLAYING_METHODS: + method = getattr(Branch, name) + if not getattr(method, REPLAY_MARK, False): + setattr(Branch, name, _marking_a_replay(method)) + + +def _marking_a_replay(method): + """Return *method*, wrapped so that the context that runs it is in a replay until it returns or raises.""" + + @functools.wraps(method) + def replay(*args, **kwargs): + token = _replaying.set(True) + try: + return method(*args, **kwargs) + finally: + _replaying.reset(token) + + setattr(replay, REPLAY_MARK, True) + return replay + + +def replay_in_progress() -> bool: + """Return whether a netbox-branching merge, revert or sync runs in this context.""" + return _replaying.get() + + +def branch_identity() -> str | None: + """Return the schema ID of the active branch, or None on main and when netbox-branching is not installed.""" + if not apps.is_installed(APP_LABEL): + return None + from netbox_branching.contextvars import active_branch + + branch = active_branch.get() + return None if branch is None else branch.schema_id + + +def activate_on(request, schema_id: str | None) -> None: + """Put netbox-branching's branch cookie for *schema_id* on *request*; without either, the request stays on main.""" + if schema_id is None or not apps.is_installed(APP_LABEL): + return + from netbox_branching.constants import COOKIE_NAME + + request.COOKIES[COOKIE_NAME] = schema_id diff --git a/netbox_interface_name_rules/family/__init__.py b/netbox_interface_name_rules/family/__init__.py index 3a78933a..53f5a25e 100644 --- a/netbox_interface_name_rules/family/__init__.py +++ b/netbox_interface_name_rules/family/__init__.py @@ -55,6 +55,7 @@ plan_installed_flat_families, plan_interface_rename, ) +from .names import ChannelReconciliationError from .previous import PreviousForms, names_the_previous_rule_gave from .prospective import ( ProspectiveInterface, @@ -83,6 +84,7 @@ __all__ = ( "UNCLAIMED_BASE_REASON", "BatchOutcome", + "ChannelReconciliationError", "ConversionCandidate", "ConversionMember", "ConversionPlan", diff --git a/netbox_interface_name_rules/family/conversion.py b/netbox_interface_name_rules/family/conversion.py index b1ed018f..4b62b7af 100644 --- a/netbox_interface_name_rules/family/conversion.py +++ b/netbox_interface_name_rules/family/conversion.py @@ -17,7 +17,7 @@ from dcim.choices import InterfaceTypeChoices from dcim.models import Interface from django.core.exceptions import ValidationError -from django.db import IntegrityError, transaction +from django.db import IntegrityError from ..naming import build_variables from ..transactions import atomic_with_events @@ -343,7 +343,7 @@ def _convert(plan, commit): # pragma: no cover - requires channelization suppor rewrite back, so a family is never half converted and a scan writes nothing at all. """ try: - with atomic_with_events(): + with atomic_with_events() as block: live = _locked_family(plan) if _is_stale(plan, live): return _refused(plan, FamilyStatus.STALE, STALE_REASON) @@ -352,7 +352,7 @@ def _convert(plan, commit): # pragma: no cover - requires channelization suppor return _refused(plan, FamilyStatus.BLOCKED, reason) members = _rewrite(plan, live) if not commit: - transaction.set_rollback(True) + block.set_rollback() except ValidationError as error: return _refused(plan, FamilyStatus.BLOCKED, " ".join(error.messages)) except IntegrityError as error: diff --git a/netbox_interface_name_rules/family/names.py b/netbox_interface_name_rules/family/names.py index 927d87fb..9450441e 100644 --- a/netbox_interface_name_rules/family/names.py +++ b/netbox_interface_name_rules/family/names.py @@ -6,9 +6,9 @@ from dcim.models import Interface from django.core.exceptions import ValidationError -from django.db import IntegrityError, transaction +from django.db import DatabaseError, IntegrityError -from ..transactions import atomic_with_events +from ..transactions import atomic_with_events, on_commit logger = logging.getLogger(__name__) @@ -38,8 +38,28 @@ def is_name_collision(error: IntegrityError) -> bool: return getattr(diagnostics, "constraint_name", None) == INTERFACE_NAME_CONSTRAINT +class ChannelReconciliationError(RuntimeError): + """The names that NetBox's parent cascade gave the channels this plugin kept could not be restored.""" + + def restore_deferred_channel_names(reconciliations): - """Restore plugin-owned names that NetBox's parent cascade changed after commit.""" + """Restore plugin-owned names that NetBox's parent cascade changed after commit. + + A database failure rolls the whole restore back and raises ``ChannelReconciliationError``, which + names each channel and the name it kept, so that an operator can rename it back. + """ + try: + _restore_channel_names(reconciliations) + except DatabaseError as error: + kept = ", ".join(f"`{cascade_name}` to `{final_name}`" for _pk, final_name, cascade_name in reconciliations) + raise ChannelReconciliationError( + f"NetBox's parent cascade renamed channels that kept their names, and restoring those names failed " + f"with {type(error).__name__}. Rename each channel back: {kept}" + ) from error + + +def _restore_channel_names(reconciliations): + """Restore each kept name in one block, and leave a channel that changed since the cascade alone.""" child_pks = [child_pk for child_pk, _final_name, _cascade_name in reconciliations] with atomic_with_events(): children = ( @@ -87,7 +107,7 @@ def reconcile_after_parent_cascade(parent_before, parent_after, channels): """Schedule restoration of the channel names NetBox's deferred parent cascade will overwrite. *channels* carries ``(child_pk, channel_id, final_name)`` for every channel the caller settled. - Registration happens on the caller's open transaction so the callback runs after NetBox's own. + The callback runs after the open transactions commit on both connections, so after NetBox's cascade. """ if parent_after == parent_before: return @@ -98,4 +118,4 @@ def reconcile_after_parent_cascade(parent_before, parent_after, channels): ) if not reconciliations: return - transaction.on_commit(lambda: restore_deferred_channel_names(reconciliations)) + on_commit(lambda: restore_deferred_channel_names(reconciliations)) diff --git a/netbox_interface_name_rules/family/template_names.py b/netbox_interface_name_rules/family/template_names.py index f98d16bf..dbd26b1f 100644 --- a/netbox_interface_name_rules/family/template_names.py +++ b/netbox_interface_name_rules/family/template_names.py @@ -2,8 +2,10 @@ # Copyright (C) 2025 Marcin Zieba """Resolve current and historical NetBox interface-template names. -The refetch and template queries here use the default manager, and the block cache is keyed by -primary key alone. Both hold because NetBox configures one database alias and no router. +The refetch and template queries here use the default manager, so NetBox's router sends them to the +active netbox-branching branch, or to main. The block cache is keyed by primary key alone. That holds +because every read of one block goes to one alias: netbox-branching activates one branch for a whole +request or job, and the plugin does not change the active branch inside a block. """ import contextlib diff --git a/netbox_interface_name_rules/jobs.py b/netbox_interface_name_rules/jobs.py index d08589ee..49093737 100644 --- a/netbox_interface_name_rules/jobs.py +++ b/netbox_interface_name_rules/jobs.py @@ -12,20 +12,29 @@ from netbox.jobs import JobRunner from netbox.registry import registry +from . import branching +from .transactions import write_alias, write_scope + + +def rule_job_kwargs(rule_id): + """Return the kwargs of a job on the rule *rule_id* in the branch of the current request, derived on the server.""" + return {"rule_id": rule_id, "branch_schema_id": branching.branch_identity(), "expected_alias": write_alias()} + def _run_under_request_processors(request, body): - """Call *body* inside every registered request processor, the way NetBox runs a script job.""" + """Call *body* inside every registered request processor; unlike NetBox, a processor that fails to enter raises.""" with ExitStack() as stack: for request_processor in registry["request_processors"]: stack.enter_context(request_processor(request)) return body() -def run_as_job_user(job, body): +def run_as_job_user(job, body, *, branch_schema_id): """Call *body*, which takes no arguments, as a request of the user who enqueued *job*. NetBox writes a change log only while a request is current, and a worker has none. The request - carries the job ID, so the change log lists the changes of one job under one request ID. + carries the job ID, so the change log lists the changes of one job under one request ID. The + request is in the branch *branch_schema_id*, or on main for None. """ if job.user is None: raise ValueError(f"Job {job.pk} has no user, and the change log must name the user of each change.") @@ -34,12 +43,16 @@ def run_as_job_user(job, body): request.method = "POST" request.user = job.user request.id = job.job_id + branching.activate_on(request, branch_schema_id) # The copy discards what the processors set; before NetBox 4.7, event_tracking keeps it when the body raises. return contextvars.copy_context().run(_run_under_request_processors, request, body) class RuleJobRunner(JobRunner): - """A background job over the rule named by rule_id, run as a request of the user who enqueued it.""" + """A background job over the rule named by rule_id, run as a request of the user who enqueued it. + + Its kwargs are those of ``rule_job_kwargs``. A job enqueued without them fails. + """ def __init__(self, job): super().__init__(job) @@ -47,25 +60,28 @@ def __init__(self, job): if not hasattr(self, "logger"): self.logger = logging.getLogger(f"netbox.jobs.{type(self).__name__}") - def run(self, *args, **kwargs): - """Run the job on the rule named by rule_id in kwargs; a missing rule is a warning, not an error.""" - run_as_job_user(self.job, functools.partial(self._run_on_rule_id, kwargs.get("rule_id"))) + def run(self, *, rule_id, branch_schema_id, expected_alias): + """Run the job on the rule *rule_id* in the branch it was enqueued from; a missing rule is a warning.""" + run_as_job_user( + self.job, + functools.partial(self._run_on_rule_id, rule_id, expected_alias), + branch_schema_id=branch_schema_id, + ) - def _run_on_rule_id(self, rule_id): + def _run_on_rule_id(self, rule_id, expected_alias): from .models import InterfaceNameRule - if not rule_id: - self.logger.warning("%s called without rule_id; skipping.", type(self).__name__) - return - rule = InterfaceNameRule.objects.filter(pk=rule_id).first() - if rule is None: - self.logger.warning("InterfaceNameRule with pk=%s does not exist; skipping.", rule_id) - return - try: - self.run_on_rule(rule) - except Exception: - self.logger.exception("%s failed on rule '%s'", self.name, rule_id) - raise + # A branch that is not ready does not activate, so the scope raises before a read on main. + with write_scope(expected_alias=expected_alias): + rule = InterfaceNameRule.objects.filter(pk=rule_id).first() + if rule is None: + self.logger.warning("InterfaceNameRule with pk=%s does not exist; skipping.", rule_id) + return + try: + self.run_on_rule(rule) + except Exception: + self.logger.exception("%s failed on rule '%s'", self.name, rule_id) + raise @abstractmethod def run_on_rule(self, rule): diff --git a/netbox_interface_name_rules/models.py b/netbox_interface_name_rules/models.py index 835f55ba..e974c291 100644 --- a/netbox_interface_name_rules/models.py +++ b/netbox_interface_name_rules/models.py @@ -9,6 +9,7 @@ from netbox.models import NetBoxModel from taggit.managers import TaggableManager +from .branching import replay_in_progress from .choices import BreakoutModeChoices from .name_template import validate_rule from .regex_safety import compile_module_type_pattern @@ -347,6 +348,9 @@ class Meta: def save(self, **kwargs): """Normalise the mode fields and validate topology and templates before a plain ORM write.""" using = kwargs.get("using") or router.db_for_write(self.__class__, instance=self) + # A merge started in an active branch replays its changes on default. + if not replay_in_progress() and using != (routed := router.db_for_write(self.__class__)): + raise RuntimeError(f"A rule save writes to {using!r}, but the router gives {routed!r} for a rule.") update_fields = kwargs.get("update_fields") if update_fields is not None: # Django accepts any iterable. Reading a generator here would leave Django an empty @@ -364,7 +368,10 @@ def save(self, **kwargs): if written := _RULE_VALIDATION_FIELDS.intersection(update_fields): # Validate the stored row on the alias Model.save() writes to, locked against a concurrent save. kwargs["using"] = using - with atomic_with_events(using=using): + with atomic_with_events() as block: + # netbox-branching routes an exempted rule model to default, which is an alias of every scope. + if using not in block.aliases: + raise RuntimeError(f"A rule save writes to {using!r}, outside the write scope {block.aliases}.") stored = self.__class__._base_manager.using(using).select_for_update().filter(pk=self.pk) # No ordering: the default one joins a nullable module type, which FOR UPDATE refuses. row = stored.order_by().values("pk", *(_RULE_VALIDATION_FIELDS - written)).first() diff --git a/netbox_interface_name_rules/rename_triggers.py b/netbox_interface_name_rules/rename_triggers.py index a21f0f6f..fcedf99d 100644 --- a/netbox_interface_name_rules/rename_triggers.py +++ b/netbox_interface_name_rules/rename_triggers.py @@ -2,15 +2,18 @@ # Copyright (C) 2025 Marcin Zieba """Rename triggers: read the previous state, decide, and reapply the rules after commit. -The receivers in ``signals.py`` pass every module, module bay and device save here. ``before_save`` -reads the previous state; ``after_save`` compares it with the saved values and, when the save is a -rename trigger, adds the trigger to the reapply plan of the transaction. Each trigger is a committed -callback, which runs only when its savepoint commits. Each trigger also appends a runner of the plan -that no savepoint rollback drops, and only the newest runner runs the plan, after every trigger. So -the plan acts on the triggers that committed, and it moves no committed callback. The plan reapplies -each module and each device at most once, from the earliest previous state of the transaction, so -it acts on the net change. A reapply that leaves an interface unrenamed, or fails, writes one journal -entry. A reapply that cannot read the committed rows is logged only. +The receivers in ``signals.py`` pass every module, module bay and device save here, with the alias of +the save, which must be the write alias. ``before_save`` reads the previous state; ``after_save`` +compares it with the saved values and, when the save is a rename trigger, adds the trigger to the +reapply plan of the transaction on the connection of the save. Each trigger is a committed callback +of that connection, which runs only when its savepoint commits. Each trigger also appends a runner +of the plan that no savepoint rollback drops, and only the newest runner runs the plan, after every +trigger. So the plan acts on the triggers that committed, and it moves no committed callback. In a +netbox-branching branch the plan therefore runs after the branch transaction commits, when a new +module has its interfaces. The plan runs in a write scope on its alias. It reapplies each module and +each device at most once, from the earliest previous state of the transaction, so it acts on the net +change. A reapply that leaves an interface unrenamed, or fails, writes one journal entry. A reapply +that cannot read the committed rows is logged only. A save that moves a module, or that changes what the names in an occupied bay are built from, also reads before the save what named the interfaces of that module and of every module nested in it. @@ -35,10 +38,11 @@ from django.db import transaction from netbox.context import current_request +from .branching import replay_in_progress from .naming import bay_naming_values, chassis_position from .rename_outcomes import OutcomeKind, RenameOutcome, renamed_count from .rule_selection import parent_type_scopes_a_rule -from .transactions import atomic_with_events +from .transactions import atomic_with_events, write_alias, write_scope if TYPE_CHECKING: from .engine import ModuleNaming @@ -139,8 +143,9 @@ class DeviceTrigger(_Trigger): @dataclasses.dataclass(eq=False) class ReapplyPlan: - """What the rename triggers of one transaction ask for, in save order, and the newest runner.""" + """What the rename triggers of one transaction on *alias* ask for, in save order, and the newest runner.""" + alias: str triggers: list = dataclasses.field(default_factory=list) runner: "PlanRunner | None" = None started: bool = dataclasses.field(default=False, init=False) # captureOnCommitCallbacks keeps run callbacks @@ -153,11 +158,15 @@ class PlanRunner: plan: ReapplyPlan def __call__(self): - """Reapply the rules for the triggers whose savepoints committed, when this is the newest runner.""" + """Reapply the rules for the triggers whose savepoints committed, when this is the newest runner. + + The reapply runs in a write scope on the plan's alias. The scope raises when the write alias differs. + """ if self.plan.runner is not self or self.plan.started: return self.plan.started = True - reapply([trigger for trigger in self.plan.triggers if trigger.kept]) + with write_scope(expected_alias=self.plan.alias): + reapply([trigger for trigger in self.plan.triggers if trigger.kept]) def reapply(triggers): @@ -653,8 +662,22 @@ def _read_bay(bay): _previous_states = {} -def before_save(sender, instance): - """Read the previous state of *instance* and hold it for its post_save. A read error propagates.""" +def _check_write_alias(using): + """Raise when a save writes through *using*, which is not the alias that the plugin writes to.""" + alias = write_alias() + if using != alias: + raise RuntimeError(f"A rename trigger save writes to {using!r}, but the write alias is {alias!r}.") + + +def before_save(sender, instance, using): + """Read the previous state of *instance*, saved through *using*, and hold it for its post_save. + + A read error propagates, and so does a save through another alias than the write alias. A save that + netbox-branching replays returns before the write-alias check: a merge writes ``default`` while a branch is active. + """ + if replay_in_progress(): + return + _check_write_alias(using) read, _ = _TRIGGERS[sender._meta.label] previous = None if instance.pk is None else read(instance) key = id(instance) @@ -667,8 +690,14 @@ def forget(reference): _previous_states[key] = (weakref.ref(instance, forget), previous) -def after_save(sender, instance, created): - """Add the save of *instance* to the reapply plan of the transaction when it is a rename trigger.""" +def after_save(sender, instance, created, using): + """Add the save of *instance* through *using* to the reapply plan of that connection when it is a rename trigger. + + A save that netbox-branching replays returns before the write-alias check, as in ``before_save``. + """ + if replay_in_progress(): + return + _check_write_alias(using) _, trigger_of = _TRIGGERS[sender._meta.label] # No entry: NetBox sent this post_save by hand, without a model save. _, previous = _previous_states.pop(id(instance), (None, None)) @@ -676,22 +705,15 @@ def after_save(sender, instance, created): if trigger is None: return trigger.author = _request_user() - connection = transaction.get_connection() + connection = transaction.get_connection(using) entries = connection.run_on_commit if connection.in_atomic_block else () - plan = ( - next( - ( - callback.plan - for _, callback, _ in entries - if isinstance(callback, PlanRunner) and not callback.plan.started - ), - None, - ) - or ReapplyPlan() - ) + plan = next( + (callback.plan for _, callback, _ in entries if isinstance(callback, PlanRunner) and not callback.plan.started), + None, + ) or ReapplyPlan(using) plan.triggers.append(trigger) # Django drops the callbacks of a rolled-back savepoint from run_on_commit, so a dropped trigger is not kept. - transaction.on_commit(trigger) + transaction.on_commit(trigger, using=using) plan.runner = PlanRunner(plan) if connection.in_atomic_block: # Without a savepoint tag no rollback but the transaction's drops the runner; it runs after every trigger. diff --git a/netbox_interface_name_rules/rule_selection.py b/netbox_interface_name_rules/rule_selection.py index fcd9c928..f515a5ac 100644 --- a/netbox_interface_name_rules/rule_selection.py +++ b/netbox_interface_name_rules/rule_selection.py @@ -4,6 +4,7 @@ import contextlib import threading +from dataclasses import dataclass from django.core.exceptions import ValidationError from django.db.models import Aggregate, F, TextField, Value @@ -11,9 +12,25 @@ from .regex_safety import compile_module_type_pattern -# Publish each loaded rule set as one new dictionary. Concurrent readers then see -# one complete version rather than a mixture of cache entries from two versions. -_RULE_CACHE = {"version": None, "exact": (), "regex": (), "memo": {}} + +@dataclass(frozen=True, slots=True) +class _RuleSnapshot: + """One enabled-rule set read from one alias; only the memo changes after publication.""" + + alias: str | None + version: str | None + exact: tuple + regex: tuple + memo: dict + + +def _empty_rule_cache(): + """Return a new snapshot that matches no alias, so the next read reloads.""" + return _RuleSnapshot(alias=None, version=None, exact=(), regex=(), memo={}) + + +# A reload rebinds this to a new snapshot, so a reader never mixes entries from two versions. +_RULE_CACHE = _empty_rule_cache() # Bound the number of module and scope contexts retained for one rule-set version. _MEMO_MAX = 4096 @@ -28,8 +45,9 @@ def pinned_rule_cache(): """Pin one enabled-rule snapshot for all selections inside the block. The first selection loads and fingerprints the rule set. Later selections in - the same thread skip the fingerprint query. Nested blocks share the snapshot. - The pin is thread-local, and an empty block does not load rules. + the same thread and from the same read alias skip the fingerprint query. + Nested blocks share the snapshot. The pin is thread-local, and an empty block + does not load rules. """ depth = getattr(_pin, "depth", 0) _pin.depth = depth + 1 @@ -41,7 +59,7 @@ def pinned_rule_cache(): _pin.depth -= 1 if _pin.depth == 0: _pin.primed = False - for attr in ("exact", "regex", "memo"): + for attr in ("alias", "exact", "regex", "memo"): _pin.__dict__.pop(attr, None) @@ -125,8 +143,8 @@ def _version_row_signature(): _ROW_SIGNATURE = _version_row_signature() -def _enabled_rules_version(): - """Return a deterministic content fingerprint of all enabled rules. +def _enabled_rules_version(alias): + """Return a deterministic content fingerprint of all enabled rules on *alias*. PostgreSQL hashes the matching and output columns in primary-key order. Each value is length-prefixed, so arbitrary text cannot create field or row boundary @@ -134,52 +152,59 @@ def _enabled_rules_version(): """ from .models import InterfaceNameRule - return InterfaceNameRule.objects.filter(enabled=True).aggregate( - fingerprint=Coalesce( - _Md5OrderedStringAgg(_ROW_SIGNATURE, Value("", output_field=TextField())), - Value("", output_field=TextField()), - ) - )["fingerprint"] + return ( + InterfaceNameRule.objects.using(alias) + .filter(enabled=True) + .aggregate( + fingerprint=Coalesce( + _Md5OrderedStringAgg(_ROW_SIGNATURE, Value("", output_field=TextField())), + Value("", output_field=TextField()), + ) + )["fingerprint"] + ) def _get_enabled_rules(): - """Return the exact rules, regex rules, and memo for the current version. + """Return the exact rules, regex rules, and memo for the current version on the read alias. Exact rules retain model ordering, which reduces to primary-key order for one module type. Regex rules are compiled once and ordered by decreasing pattern - length, then primary key. A reload publishes one new cache dictionary so a + length, then primary key. A reload publishes one new cache snapshot so a concurrent reader cannot combine values from two versions. """ # One module-level cache, replaced atomically. global _RULE_CACHE # noqa: PLW0603 + # A branch can hold other rules than main, so a snapshot belongs to the alias it was read from. + alias = _enabled_module_rules().db pinned = getattr(_pin, "depth", 0) > 0 - if pinned and getattr(_pin, "primed", False): + if pinned and getattr(_pin, "primed", False) and _pin.alias == alias: # Return the thread's snapshot. Another thread can replace the shared cache. return _pin.exact, _pin.regex, _pin.memo cache = _RULE_CACHE - version = _enabled_rules_version() - if cache["version"] != version: - rules = list(_enabled_module_rules().order_by("module_type__model", "pk")) + version = _enabled_rules_version(alias) + if (cache.alias, cache.version) != (alias, version): + rules = list(_enabled_module_rules().using(alias).order_by("module_type__model", "pk")) exact = tuple(rule for rule in rules if not rule.module_type_is_regex) regex_rules = sorted( (rule for rule in rules if rule.module_type_is_regex), key=lambda rule: (-len(rule.module_type_pattern or ""), rule.pk), ) regex = tuple((compile_stored_pattern(rule.module_type_pattern), rule) for rule in regex_rules) - cache = {"version": version, "exact": exact, "regex": regex, "memo": {}} + cache = _RuleSnapshot(alias=alias, version=version, exact=exact, regex=regex, memo={}) _RULE_CACHE = cache if pinned: # Keep a private memo so another thread cannot clear this batch's entries. - _pin.exact = cache["exact"] - _pin.regex = cache["regex"] - _pin.memo = dict(cache["memo"]) + _pin.alias = alias + _pin.exact = cache.exact + _pin.regex = cache.regex + _pin.memo = dict(cache.memo) _pin.primed = True return _pin.exact, _pin.regex, _pin.memo - return cache["exact"], cache["regex"], cache["memo"] + return cache.exact, cache.regex, cache.memo def _scope_ids(parent_module_type, device_type, platform): diff --git a/netbox_interface_name_rules/signals.py b/netbox_interface_name_rules/signals.py index 8748ac8d..ebcf0f48 100644 --- a/netbox_interface_name_rules/signals.py +++ b/netbox_interface_name_rules/signals.py @@ -1,6 +1,9 @@ # SPDX-License-Identifier: Apache-2.0 # Copyright (C) 2025 Marcin Zieba -"""Django receivers: module, module bay and device saves go to the rename-trigger lifecycle.""" +"""Django receivers: module, module bay and device saves go to the rename-trigger lifecycle with their alias. + +The receivers do not read ``raw``: a raw save, such as a fixture load, is a rename trigger too. +""" import logging @@ -15,39 +18,39 @@ @receiver(pre_save, sender="dcim.Module", dispatch_uid="interface_name_rules_pre_save_module") -def on_module_pre_save(sender, instance, **kwargs): +def on_module_pre_save(sender, instance, using, **kwargs): """Pass a module save to the rename-trigger lifecycle before the row is written.""" - rename_triggers.before_save(sender, instance) + rename_triggers.before_save(sender, instance, using) @receiver(post_save, sender="dcim.Module", dispatch_uid="interface_name_rules_post_save_module") -def on_module_saved(sender, instance, created, **kwargs): +def on_module_saved(sender, instance, created, using, **kwargs): """Pass a module save to the rename-trigger lifecycle after the row is written.""" - rename_triggers.after_save(sender, instance, created) + rename_triggers.after_save(sender, instance, created, using) @receiver(pre_save, sender="dcim.ModuleBay", dispatch_uid="interface_name_rules_pre_save_module_bay") -def on_module_bay_pre_save(sender, instance, **kwargs): +def on_module_bay_pre_save(sender, instance, using, **kwargs): """Pass a module bay save to the rename-trigger lifecycle before the row is written.""" - rename_triggers.before_save(sender, instance) + rename_triggers.before_save(sender, instance, using) @receiver(post_save, sender="dcim.ModuleBay", dispatch_uid="interface_name_rules_post_save_module_bay") -def on_module_bay_saved(sender, instance, created, **kwargs): +def on_module_bay_saved(sender, instance, created, using, **kwargs): """Pass a module bay save to the rename-trigger lifecycle after the row is written.""" - rename_triggers.after_save(sender, instance, created) + rename_triggers.after_save(sender, instance, created, using) @receiver(pre_save, sender="dcim.Device", dispatch_uid="interface_name_rules_pre_save_device") -def on_device_pre_save(sender, instance, **kwargs): +def on_device_pre_save(sender, instance, using, **kwargs): """Pass a device save to the rename-trigger lifecycle before the row is written.""" - rename_triggers.before_save(sender, instance) + rename_triggers.before_save(sender, instance, using) @receiver(post_save, sender="dcim.Device", dispatch_uid="interface_name_rules_post_save_device") -def on_device_saved(sender, instance, created, **kwargs): +def on_device_saved(sender, instance, created, using, **kwargs): """Pass a device save to the rename-trigger lifecycle after the row is written.""" - rename_triggers.after_save(sender, instance, created) + rename_triggers.after_save(sender, instance, created, using) # --------------------------------------------------------------------------- diff --git a/netbox_interface_name_rules/tests/branch_cases.py b/netbox_interface_name_rules/tests/branch_cases.py new file mode 100644 index 00000000..7a596367 --- /dev/null +++ b/netbox_interface_name_rules/tests/branch_cases.py @@ -0,0 +1,236 @@ +# SPDX-License-Identifier: Apache-2.0 +# Copyright (C) 2025 Marcin Zieba +"""Test cases that provision a real netbox-branching branch, and the rows and values that the branch tests share. + +The cases skip when netbox-branching is not installed. This module imports netbox-branching only inside +functions, so it imports where netbox-branching is not installed. +""" + +from unittest import skipUnless + +from core.models import ObjectChange +from dcim.models import Interface, InterfaceTemplate, Module, ModuleBay +from django.apps import apps +from django.contrib.auth import get_user_model +from django.contrib.contenttypes.models import ContentType +from django.db import connections, transaction +from django.test import TransactionTestCase +from django.urls import reverse + +from netbox_interface_name_rules.engine import apply_rule_to_existing, supports_channelization +from netbox_interface_name_rules.models import InterfaceNameRule +from netbox_interface_name_rules.tests.helpers import ( + CHANNELIZED, + FLAT, + PLAIN_TYPE, + REQUIRES_CHANNELIZATION, + activate, + branch_cookie, + channelized_module_type, + make_device, + make_device_type, + make_manufacturer, + make_module_bay_templates, + make_module_type, +) + +User = get_user_model() +BRANCHING_INSTALLED = apps.is_installed("netbox_branching") +BRANCHING_SKIP_REASON = "netbox-branching is not installed" +# Each connection starts from its own value, so a test can tell which value came back where. +DEFAULT_BEFORE = "3s" +BRANCH_BEFORE = "4s" +# The value that PostgreSQL gives a new session. +SERVER_DEFAULT = "0" +# The names of the flat family that each ConversionCase builds in bay 3. +FLAT_NAMES = ("xe-0/0/3:0", "xe-0/0/3:1", "xe-0/0/3:2", "xe-0/0/3:3") + + +def remove_branch(branch): + """Close the connection of *branch*, then drop its schema.""" + connections[branch.connection_name].close() + branch.deprovision() + + +@skipUnless(BRANCHING_INSTALLED, BRANCHING_SKIP_REASON) +class BranchTestCase(TransactionTestCase): + """Provision real branches. Each branch is removed when its test ends, because ``--reuse-db`` keeps schemas.""" + + def provision_branch(self, name, user): + """Return a new branch named *name*, provisioned by *user* as netbox-branching's own tests do.""" + from netbox_branching.models import Branch + + branch = Branch(name=name) + branch.save(provision=False) + self.addCleanup(remove_branch, branch) + branch.provision(user=user) + # provision() writes the status with a queryset update, which the instance does not see. + branch.refresh_from_db() + return branch + + +class BranchWriteCase(BranchTestCase): + """A superuser logged in with netbox-branching's cookie of a branch provisioned from the rows ``build`` made.""" + + PREFIX = "" + + def setUp(self): + self.user = User.objects.create_superuser(username=f"{self.PREFIX.lower()}-operator") + self.build() + self.branch = self.provision_branch(self.PREFIX, self.user) + self.alias = self.branch.connection_name + self.client.force_login(self.user) + self.client.cookies[branch_cookie()] = self.branch.schema_id + + def build(self): + """Create the rows on main that the branch copies.""" + raise NotImplementedError + + def in_branch(self): + return activate(self.branch) + + def bay(self, position): + return ModuleBay.objects.get(device=self.device, name=f"Bay {position}") + + def interfaces_at(self, position): + """Return ``(pk, name)`` of each interface of the module in the bay at *position*, on the active branch or main.""" + interfaces = Interface.objects.filter(device=self.device, module__module_bay__name=f"Bay {position}") + return list(interfaces.values_list("pk", "name")) + + def interfaces_in_branch(self, position): + """Return ``(pk, name)`` of each interface in the branch of the module in the bay at *position*.""" + with self.in_branch(): + return self.interfaces_at(position) + + def names_in_branch(self, position): + """Return the sorted interface names in the branch of the module in the bay at *position*.""" + return sorted(name for _, name in self.interfaces_in_branch(position)) + + def apply_url(self, rule): + return reverse("plugins:netbox_interface_name_rules:interfacenamerule_apply_detail", kwargs={"pk": rule.pk}) + + def change_diffs(self): + """Return the ChangeDiff rows of the branch, which netbox-branching keeps on ``default``.""" + from netbox_branching.models import ChangeDiff + + return ChangeDiff.objects.filter(branch=self.branch) + + def branch_updates_of(self, instance): + """Return ``(name before, name after)`` for each update record of *instance* in the branch.""" + changes = ObjectChange.objects.using(self.alias).filter( + changed_object_type=ContentType.objects.get_for_model(instance), changed_object_id=instance.pk + ) + return [(change.prechange_data["name"], change.postchange_data["name"]) for change in changes] + + +class PlainModuleCase(BranchWriteCase): + """One device with one module bay, and a module whose one interface NetBox named ``0``.""" + + def build(self): + manufacturer = make_manufacturer(self.PREFIX) + device_type = make_device_type(manufacturer, self.PREFIX) + make_module_bay_templates(device_type, ("Bay 0",)) + self.device = make_device(self.PREFIX, device_type) + self.device_type = device_type + self.module_type = make_module_type(manufacturer, self.PREFIX) + InterfaceTemplate.objects.create(module_type=self.module_type, name="{module}", type=PLAIN_TYPE) + # No rule exists yet, so the interface keeps NetBox's raw name. + self.module = Module.objects.create(device=self.device, module_bay=self.bay(0), module_type=self.module_type) + self.interface = Interface.objects.get(module=self.module) + + +@skipUnless(supports_channelization(), REQUIRES_CHANNELIZATION) +class ConversionCase(BranchWriteCase): + """A flat family built on main by a flat rule, and the rule then switched to the channelized topology.""" + + def build(self): + manufacturer = make_manufacturer(self.PREFIX) + device_type = make_device_type(manufacturer, self.PREFIX) + make_module_bay_templates(device_type, ("Bay 0", "Bay 1", "Bay 2", "Bay 3")) + self.device = make_device(self.PREFIX, device_type) + module_type = make_module_type(manufacturer, self.PREFIX) + InterfaceTemplate.objects.create(module_type=module_type, name="{module}", type=PLAIN_TYPE) + self.rule = InterfaceNameRule.objects.create( + module_type=module_type, + name_template="xe-0/0/{bay_position}:{channel}", + breakout_mode=FLAT, + channel_count=4, + channel_start=0, + ) + # The rename trigger builds the family when the install commits. + with transaction.atomic(): + self.module = Module.objects.create(device=self.device, module_bay=self.bay(3), module_type=module_type) + self.rule.snapshot() + self.rule.breakout_mode = CHANNELIZED + self.rule.parent_name_template = "et-0/0/{bay_position}" + self.rule.save() + self.base = Interface.objects.get(module=self.module, name="xe-0/0/3:0") + + +@skipUnless(supports_channelization(), REQUIRES_CHANNELIZATION) +class ChannelCase(BranchWriteCase): + """Empty module bays, a channelized module type, and an interface on the target of channel 2 at each position. + + NetBox names a family ```` and ``:``. The rule that ``add_rule`` creates + keeps channel 2 at its old name while it renames the parent. NetBox's cascade then renames the kept + channel after the parent, and the plugin's reconciliation gives it back the name it kept. + """ + + POSITIONS = ("1",) + + def build(self): + manufacturer = make_manufacturer(self.PREFIX) + device_type = make_device_type(manufacturer, self.PREFIX) + make_module_bay_templates(device_type, [f"Bay {index}" for index in range(max(map(int, self.POSITIONS)) + 1)]) + self.device = make_device(self.PREFIX, device_type) + self.module_type = channelized_module_type(manufacturer, f"{self.PREFIX}-QSFP") + for position in self.POSITIONS: + Interface.objects.create(device=self.device, name=f"xe-0/0/{position}:1", type=PLAIN_TYPE) + + def add_rule(self): + self.rule = InterfaceNameRule.objects.create( + module_type=self.module_type, + name_template="xe-0/0/{bay_position}:{channel}", + parent_name_template="et-0/0/{bay_position}", + breakout_mode=CHANNELIZED, + channel_count=4, + channel_start=0, + ) + + @staticmethod + def kept(position): + """Return the names of the family at *position* after the reconciliation gave channel 2 its kept name.""" + return [ + f"{position}:2", + f"et-0/0/{position}", + f"xe-0/0/{position}:0", + f"xe-0/0/{position}:2", + f"xe-0/0/{position}:3", + ] + + @staticmethod + def cascaded(position): + """Return the names of the family at *position* when channel 2 carries the name of NetBox's cascade.""" + return sorted([f"et-0/0/{position}:2", f"et-0/0/{position}", *(f"xe-0/0/{position}:{c}" for c in (0, 2, 3))]) + + +class KeptChannelCase(ChannelCase): + """The families installed on main before the rule exists, so they keep NetBox's raw names, and the rule.""" + + def build(self): + super().build() + self.modules = { + position: Module.objects.create( + device=self.device, module_bay=self.bay(position), module_type=self.module_type + ) + for position in self.POSITIONS + } + self.add_rule() + + def parent(self, position): + return Interface.objects.get(module=self.modules[position], channels__isnull=False) + + def apply(self, position): + """Apply the rule in the branch to the family at *position*.""" + with self.in_branch(): + return apply_rule_to_existing(self.rule, interface_ids=[self.parent(position).pk]) diff --git a/netbox_interface_name_rules/tests/comment_blocks.json b/netbox_interface_name_rules/tests/comment_blocks.json index fd8a6516..d2a3c852 100644 --- a/netbox_interface_name_rules/tests/comment_blocks.json +++ b/netbox_interface_name_rules/tests/comment_blocks.json @@ -267,10 +267,6 @@ [ "# Skip duplicate detection when the user lacks view permission — we cannot\n# query existing rules without it, so add-only users always land on the\n# create form (potentially allowing duplicates).", 1 - ], - [ - "# instance is intentionally omitted: InterfaceNameRule does not\n# inherit JobsMixin, so passing instance= would fail full_clean().\n# The job is still named and findable in Core → Jobs.", - 1 ] ] } diff --git a/netbox_interface_name_rules/tests/helpers.py b/netbox_interface_name_rules/tests/helpers.py index 143c9606..cc463f13 100644 --- a/netbox_interface_name_rules/tests/helpers.py +++ b/netbox_interface_name_rules/tests/helpers.py @@ -1,6 +1,6 @@ # SPDX-License-Identifier: Apache-2.0 # Copyright (C) 2025 Marcin Zieba -"""Builders for the objects the tests need, a runner for the background jobs, and a reader of their webhooks. +"""Builders, the values and test cases that several test modules share, a job runner and a webhook reader. Every builder takes a *prefix* and derives names and slugs from it. Test classes share one database per worker, so a class that names its objects after itself cannot collide with another class, and a @@ -14,21 +14,46 @@ import django_rq from core.events import OBJECT_CREATED, OBJECT_DELETED, OBJECT_UPDATED from core.models import Job, ObjectType +from dcim.choices import InterfaceTypeChoices from dcim.models import ( Device, DeviceRole, DeviceType, Interface, + InterfaceTemplate, Manufacturer, + Module, + ModuleBay, ModuleBayTemplate, ModuleType, Site, ) from django.contrib.auth import get_user_model +from django.db import connections +from django.http import HttpRequest +from django.test import TestCase from extras.choices import EventRuleActionChoices from extras.models import EventRule, Webhook +from netbox.context import current_request, events_queue +from rq import Worker +from rq.job import Job as RQJob +from netbox_interface_name_rules.choices import BreakoutModeChoices from netbox_interface_name_rules.models import InterfaceNameRule +from netbox_interface_name_rules.transactions import READ_LOCK_TIMEOUT, SET_LOCK_TIMEOUT + +# Resolved defensively so this module still imports on NetBox releases without channelization. +CHANNEL_TYPE = getattr(InterfaceTypeChoices, "TYPE_CHANNEL", "channel") +PARENT_TYPE = InterfaceTypeChoices.TYPE_40GE_QSFP_PLUS +PLAIN_TYPE = InterfaceTypeChoices.TYPE_10GE_SFP_PLUS +FLAT = BreakoutModeChoices.FLAT +CHANNELIZED = BreakoutModeChoices.CHANNELIZED +TEST_PASSWORD = "testpass123" # noqa: S105 - Test credential only. +PLUGIN_LOGGER = "netbox_interface_name_rules" +REQUIRES_CHANNELIZATION = "requires a NetBox that models channelized interfaces (4.7+)" +REQUIRES_NO_CHANNELIZATION = "requires a NetBox that cannot model channelized interfaces (4.6 and older)" +# On NetBox 4.5 and older the token stays literal in the interface name. +REQUIRES_VC_POSITION_TOKEN = "requires a NetBox that resolves {vc_position} in template names (4.6+)" # noqa: S105 - Skip reason, not a credential. def slug_for(prefix: str, suffix: str = "") -> str: @@ -120,6 +145,48 @@ def make_unrunnable_rule(prefix: str) -> InterfaceNameRule: return rule +def activate(branch): + """Return netbox-branching's context manager that makes *branch* active, or main for None.""" + # Imported here, so that this module imports where netbox-branching is not installed. + from netbox_branching.utilities import activate_branch + + return activate_branch(branch) + + +def branch_cookie() -> str: + """Return the name of netbox-branching's branch cookie.""" + from netbox_branching.constants import COOKIE_NAME + + return COOKIE_NAME + + +def lock_timeout(alias: str) -> str: + """Return the ``lock_timeout`` of the session of *alias*, read as the write scope reads it.""" + with connections[alias].cursor() as cursor: + cursor.execute(READ_LOCK_TIMEOUT) + return cursor.fetchone()[0] + + +def set_lock_timeout(alias: str, value: str) -> None: + """Set the ``lock_timeout`` of the session of *alias*, as the write scope sets it.""" + with connections[alias].cursor() as cursor: + cursor.execute(SET_LOCK_TIMEOUT, [value]) + + +def register_a_worker(test_case): + """Register an RQ worker of the default queue until *test_case* ends, so NetBox accepts work; no process runs it.""" + worker = Worker(["default"], connection=django_rq.get_connection()) + worker.register_birth() + test_case.addCleanup(worker.register_death) + + +def queued_job(test_case, job): + """Return the queue's record of *job*, which a worker reads from Redis, and delete it when *test_case* ends.""" + queued = RQJob.fetch(str(job.job_id), connection=django_rq.get_connection()) + test_case.addCleanup(queued.delete) + return queued + + def run_job_logged(test_case, runner, raises=None, **kwargs): """Run *runner* with *kwargs*, assert that it raises *raises* when given, and return the records it logs.""" with test_case.assertLogs(f"netbox.jobs.{type(runner).__name__}", level="INFO") as logs: @@ -157,3 +224,196 @@ def queued_webhook_jobs(event_rule) -> list: def queued_webhooks(event_rule) -> list[tuple[str, str]]: """Return the event type and object name of each webhook that *event_rule* queued, sorted.""" return sorted((job.kwargs["event_type"], job.kwargs["data"]["name"]) for job in queued_webhook_jobs(event_rule)) + + +@contextlib.contextmanager +def row_lock_in_another_session(alias): + """Yield a function that locks one interface row on *alias* from a second session until the block ends.""" + other = connections.create_connection(alias) + try: + other.set_autocommit(False) + + def lock(pk): + with other.cursor() as cursor: + cursor.execute("SELECT id FROM dcim_interface WHERE id = %s FOR UPDATE", [pk]) + + yield lock + finally: + other.rollback() + other.close() + + +@contextlib.contextmanager +def interface_signal(signal, receiver): + """Connect *receiver* to the model *signal* of Interface while the block runs.""" + signal.connect(receiver, sender=Interface, weak=False) + try: + yield + finally: + signal.disconnect(receiver, sender=Interface) + + +def build_device(prefix, bay_positions=(), **device_kwargs): + """Create a manufacturer, a device type with module bays at *bay_positions*, and one device.""" + slug = prefix.lower() + manufacturer = Manufacturer.objects.create(name=f"{prefix}Mfg", slug=f"{slug}-mfg") + device_type = DeviceType.objects.create(manufacturer=manufacturer, model=f"{prefix}-Dev", slug=f"{slug}-dev") + for position in bay_positions: + ModuleBayTemplate.objects.create(device_type=device_type, name=f"Bay {position}", position=position) + role = DeviceRole.objects.create(name=f"{prefix}Role", slug=f"{slug}-role") + site = Site.objects.create(name=f"{prefix}Site", slug=f"{slug}-site") + device = Device.objects.create(name=f"{slug}-sw1", device_type=device_type, role=role, site=site, **device_kwargs) + return manufacturer, device + + +def channelized_family(module_type, parent_name, child_names, channels=4): + """Add one channelized parent template plus its channel templates to *module_type*. + + *child_names* maps a channel_id to the template name that channel takes. + """ + parent = InterfaceTemplate.objects.create( + module_type=module_type, name=parent_name, type=PARENT_TYPE, channels=channels + ) + for channel_id, name in child_names.items(): + InterfaceTemplate.objects.create( + module_type=module_type, + name=name, + type=CHANNEL_TYPE, + parent=parent, + channel_id=channel_id, + ) + return parent + + +def channelized_module_type(manufacturer, model, channels=4, child_channel_ids=(1, 2, 3, 4), child_names=None): + """Create a ModuleType whose interface templates form a channelized family. + + The parent template is ``{module}`` with *channels* set; each entry in *child_channel_ids* adds a + channel-type template bound to it. *child_names* maps a channel_id to a template name, defaulting + to the upstream ``:`` convention. + """ + module_type = ModuleType.objects.create(manufacturer=manufacturer, model=model, part_number=model) + names = child_names or {channel_id: f"{{module}}:{channel_id}" for channel_id in child_channel_ids} + channelized_family( + module_type, "{module}", {channel_id: names[channel_id] for channel_id in child_channel_ids}, channels=channels + ) + return module_type + + +def plain_module_type(manufacturer, model, iface_type=PARENT_TYPE): + """Create a ModuleType with a single plain (non-channelized) port template.""" + module_type = ModuleType.objects.create(manufacturer=manufacturer, model=model, part_number=model) + InterfaceTemplate.objects.create(module_type=module_type, name="{module}", type=iface_type) + return module_type + + +def token_module_type(manufacturer, model, *template_names, iface_type=PLAIN_TYPE): + """Create a ModuleType whose interface templates are named *template_names*, in order.""" + module_type = ModuleType.objects.create(manufacturer=manufacturer, model=model, part_number=model) + for name in template_names: + InterfaceTemplate.objects.create(module_type=module_type, name=name, type=iface_type) + return module_type + + +def install_form(bay, module_type): + """Return the form data of NetBox's module edit view that installs *module_type* in *bay*.""" + return { + "device": bay.device_id, + "module_bay": bay.pk, + "module_type": module_type.pk, + "status": "active", + "replicate_components": "on", + } + + +def names_of(module): + """Return the sorted interface names of *module* on the active branch.""" + return sorted(Interface.objects.filter(module=module).values_list("name", flat=True)) + + +@contextlib.contextmanager +def request_context(user): + """Set the request and the event queue as NetBox's event_tracking does, without its flush.""" + request = HttpRequest() + request.user = user + request.id = uuid.uuid4() + request_token = current_request.set(request) + queue_token = events_queue.set({}) + try: + yield + finally: + events_queue.reset(queue_token) + current_request.reset(request_token) + + +class WriteInterfacesTo: + """A database router that sends each write of an interface to one alias.""" + + def __init__(self, alias): + self.alias = alias + + def db_for_write(self, model, **hints): + return self.alias if model is Interface else None + + +class ChannelizationTestCase(TestCase): + """Install helpers shared by the channelized module-install test cases.""" + + def _install(self, module_type, position, run_rules=True): + """Install a module into the bay at *position*; run the post-commit rename unless told not to. + + ``run_rules=False`` leaves the freshly instantiated (raw-named) family in place so a test can + call the engine directly and assert its return value. + """ + bay = ModuleBay.objects.get(device=self.device, name=f"Bay {position}") + if run_rules: + with self.captureOnCommitCallbacks(execute=True): + module = Module.objects.create(device=self.device, module_bay=bay, module_type=module_type) + else: + module = Module.objects.create(device=self.device, module_bay=bay, module_type=module_type) + return module, bay + + @staticmethod + def _names(module): + """Return the sorted interface names of *module*.""" + return sorted(Interface.objects.filter(module=module).values_list("name", flat=True)) + + @staticmethod + def _parent(module): + """Return the channelized parent interface of *module*.""" + return Interface.objects.get(module=module, channels__isnull=False) + + @staticmethod + def _child(module, channel_id): + """Return the channel subinterface of *module* bound to *channel_id*.""" + return Interface.objects.get(module=module, channel_id=channel_id) + + +class VcDriftTestCase(ChannelizationTestCase): + """VC transitions go through a real ``Device.save()`` so the plugin's signals do the scheduling.""" + + def _save_vc_state(self, device, virtual_chassis, position): + with self.captureOnCommitCallbacks(execute=True): + device.virtual_chassis = virtual_chassis + device.vc_position = position + device.save() + + def _join(self, vc, position, device=None): + """Add *device* to *vc* at *position* — the join direction (fallback → position).""" + self._save_vc_state(device or self.device, vc, position) + + def _renumber(self, position, device=None): + """Move *device* to another position inside its VC — the renumber direction (P → Q).""" + device = device or self.device + self._save_vc_state(device, device.virtual_chassis, position) + + def _leave(self, device=None): + """Remove *device* from its VC — the leave direction (position → fallback).""" + self._save_vc_state(device or self.device, None, None) + + def _install_on(self, device, module_type, position): + """Install a module of *module_type* into *device*'s bay at *position*, rules and all.""" + bay = ModuleBay.objects.get(device=device, name=f"Bay {position}") + with self.captureOnCommitCallbacks(execute=True): + module = Module.objects.create(device=device, module_bay=bay, module_type=module_type) + return module, bay diff --git a/netbox_interface_name_rules/tests/test_api.py b/netbox_interface_name_rules/tests/test_api.py index 8382f133..1563b24f 100644 --- a/netbox_interface_name_rules/tests/test_api.py +++ b/netbox_interface_name_rules/tests/test_api.py @@ -2,11 +2,17 @@ # Copyright (C) 2025 Marcin Zieba """Tests for the REST API endpoints.""" +from unittest import skipUnless + +from core.choices import JobStatusChoices +from core.models import Job from dcim.models import DeviceType, Manufacturer, ModuleType +from netbox import jobs as netbox_jobs from rest_framework import status from utilities.testing import APITestCase from netbox_interface_name_rules.models import InterfaceNameRule +from netbox_interface_name_rules.tests.helpers import make_manufacturer, make_module_type, register_a_worker class InterfaceNameRuleAPITest(APITestCase): @@ -176,3 +182,28 @@ def test_update_rule_to_regex(self): rule.refresh_from_db() self.assertTrue(rule.module_type_is_regex) self.assertEqual(rule.module_type_pattern, "QSFP-.*") + + +@skipUnless(hasattr(netbox_jobs, "AsyncAPIJob"), "NetBox before 4.7 runs no REST request in the background") +class BackgroundRuleRequestOnMainTest(APITestCase): + """The rule endpoints refuse a background request only in a branch; on main NetBox enqueues it.""" + + model = InterfaceNameRule + view_namespace = "plugins-api:netbox_interface_name_rules" + user_permissions = ("netbox_interface_name_rules.change_interfacenamerule",) + + def setUp(self): + super().setUp() + register_a_worker(self) + + def test_a_background_rule_request_on_main_is_enqueued(self): + rule = InterfaceNameRule.objects.create( + module_type=make_module_type(make_manufacturer("APIBackground"), "APIBackground"), + name_template="et-0/0/{bay_position}", + ) + payload = [{"id": rule.pk, "name_template": "xe-0/0/{bay_position}"}] + + response = self.client.patch(f"{self._get_list_url()}?background=true", payload, format="json", **self.header) + + self.assertEqual(response.status_code, status.HTTP_202_ACCEPTED, response.content) + self.assertEqual(Job.objects.get(pk=response.data["job"]["id"]).status, JobStatusChoices.STATUS_PENDING) diff --git a/netbox_interface_name_rules/tests/test_bay_edit_trigger.py b/netbox_interface_name_rules/tests/test_bay_edit_trigger.py index 7862ecbf..d6344611 100644 --- a/netbox_interface_name_rules/tests/test_bay_edit_trigger.py +++ b/netbox_interface_name_rules/tests/test_bay_edit_trigger.py @@ -20,64 +20,32 @@ from rest_framework import status from utilities.testing import APITestCase -from netbox_interface_name_rules.choices import BreakoutModeChoices from netbox_interface_name_rules.engine import supports_vc_position_token from netbox_interface_name_rules.family import supports_module_moves from netbox_interface_name_rules.models import InterfaceNameRule from netbox_interface_name_rules.rename_triggers import ModuleTrigger, PlanRunner +from netbox_interface_name_rules.tests.helpers import PLAIN_TYPE, REQUIRES_VC_POSITION_TOKEN from netbox_interface_name_rules.tests.out_of_band import rename_out_of_band -from netbox_interface_name_rules.tests.test_module_move_trigger import ( +from netbox_interface_name_rules.tests.trigger_cases import ( CHASSIS_RULES, - PLAIN_TYPE, REQUIRES_SUBTREE_MOVES, TAKEN, UNAVAILABLE, - ModuleMoveTestCase, - _fail_the_naming_read, - _journal, - _module_reapplies, - _MoveFixture, - _naming_reads, - _reapplied, - _reject_interface_updates, + BayEditTestCase, + MoveFixture, + fail_the_naming_read, + flat_rule, + journal, + module_reapplies, + naming_reads, + previous_state_read_fails, + reapplied, + reject_interface_updates, ) -from netbox_interface_name_rules.tests.test_rename_triggers import _previous_state_read_fails -from netbox_interface_name_rules.tests.test_vc_drift import REQUIRES_VC_POSITION_TOKEN BAY_STATE_READ = re.compile(r'SELECT "dcim_modulebay"\."position".* FROM "dcim_modulebay"') -def _flat_rule(module_type, name_template, **scope): - return InterfaceNameRule.objects.create( - module_type=module_type, - name_template=name_template, - breakout_mode=BreakoutModeChoices.FLAT, - channel_count=2, - channel_start=0, - **scope, - ) - - -class BayEditTestCase(ModuleMoveTestCase): - """Install modules and edit their bays through real saves, with the committed callbacks run.""" - - @classmethod - def setUpTestData(cls): - super().setUpTestData() - cls.plain_type = cls._module_type("Plain", "{module}") - InterfaceNameRule.objects.create(module_type=cls.plain_type, name_template="et-{vc_position}/0/{bay_position}") - - @staticmethod - def _save_edit(bay, **values): - for field, value in values.items(): - setattr(bay, field, value) - bay.save() - - def _edit(self, bay, **values): - with self.captureOnCommitCallbacks(execute=True): - self._save_edit(bay, **values) - - class BayEditTest(BayEditTestCase): """The module in an edited bay gets the names its rule gives for the bay's new position or name.""" @@ -101,7 +69,7 @@ def test_a_raw_name_whose_bay_position_brought_the_chassis_token_is_recognised_a self._edit(bay, position="2") self.assertEqual(self._names(module), ["p1/2"]) - self.assertEqual(_journal(module), []) + self.assertEqual(journal(module), []) @skipUnless(supports_vc_position_token(), REQUIRES_VC_POSITION_TOKEN) def test_a_raw_name_whose_fallback_reads_the_bay_is_recognised_after_a_bay_edit_off_a_chassis(self): @@ -116,7 +84,7 @@ def test_a_raw_name_whose_fallback_reads_the_bay_is_recognised_after_a_bay_edit_ self._edit(bay, position="4") self.assertEqual(self._names(module), ["pxe-4"]) - self.assertEqual(_journal(module), []) + self.assertEqual(journal(module), []) def test_a_position_edit_renames_the_module_for_the_new_position(self): bay = self._bay(self.device) @@ -126,7 +94,7 @@ def test_a_position_edit_renames_the_module_for_the_new_position(self): self._edit(bay, position="5") self.assertEqual(self._names(module), ["et-1/0/5"]) - self.assertEqual(_journal(module), []) + self.assertEqual(journal(module), []) def test_a_base_rule_is_renamed_from_the_raw_name_the_old_position_gave(self): bay = self._bay(self.device) @@ -145,7 +113,7 @@ def test_a_name_edit_renames_the_module_when_its_position_takes_the_number_from_ self._edit(bay, name="Bay 7") self.assertEqual(self._names(module), ["ge-1/0/7"]) - self.assertEqual(_journal(module), []) + self.assertEqual(journal(module), []) def test_a_name_edit_that_no_variable_reads_changes_nothing(self): numbered = self._install(self.plain_type, self._bay(self.device)) @@ -153,7 +121,7 @@ def test_a_name_edit_that_no_variable_reads_changes_nothing(self): named = self._install(self.fixed_type, self._bay(self.device, "Bay 3")) control = self._install(self.plain_type, self._bay(self.device, "Bay 1")) - with _naming_reads() as reads, self.captureOnCommitCallbacks(execute=True), transaction.atomic(): + with naming_reads() as reads, self.captureOnCommitCallbacks(execute=True), transaction.atomic(): self._save_edit(self._bay(self.device), name="Uplink 0") self._save_edit(self._bay(self.device, "Bay 3"), name="Slot 3") self._save_edit(self._bay(self.device, "Bay 1"), position="4") @@ -163,7 +131,7 @@ def test_a_name_edit_that_no_variable_reads_changes_nothing(self): (self._names(numbered), self._names(named), self._names(control)), (["et-1/0/0", "operator-name"], ["ge-1/0/3"], ["et-1/0/4"]), ) - self.assertEqual((_journal(numbered), _journal(named)), ([], [])) + self.assertEqual((journal(numbered), journal(named)), ([], [])) def test_only_an_occupied_bay_whose_position_or_name_changes_reads_the_naming(self): occupied = self._bay(self.device) @@ -171,8 +139,8 @@ def test_only_an_occupied_bay_whose_position_or_name_changes_reads_the_naming(se control = self._install(self.plain_type, self._bay(self.device, "Bay 1")) with ( - _naming_reads() as reads, - _module_reapplies() as reapplies, + naming_reads() as reads, + module_reapplies() as reapplies, self.captureOnCommitCallbacks(execute=True), transaction.atomic(), ): @@ -200,7 +168,7 @@ def test_a_subinterface_does_not_stop_the_rename(self): self._edit(bay, position="5") self.assertEqual(self._names(module), ["et-1/0/0.100", "et-1/0/5"]) - self.assertEqual(_journal(module), []) + self.assertEqual(journal(module), []) class NestedBayEditTest(BayEditTestCase): @@ -226,7 +194,7 @@ def test_the_module_in_the_bay_and_the_modules_nested_in_it_are_renamed_for_the_ self._edit(bay, position="2") self.assertEqual((self._names(card), self._names(optic)), (["p2-1"], ["et-1/2/1"])) - self.assertEqual(_journal(card), []) + self.assertEqual(journal(card), []) def test_the_subtree_reports_in_one_journal_entry_on_the_module_in_the_edited_bay(self): card_type = self._card_type("Two Port Card", "1") @@ -242,16 +210,16 @@ def test_the_subtree_reports_in_one_journal_entry_on_the_module_in_the_edited_ba (runner,) = [callback for callback in callbacks if isinstance(callback, PlanRunner)] self.assertEqual([trigger.pk for trigger in runner.plan.triggers], [card.pk]) - with connection.execute_wrapper(_reject_interface_updates), self.assertLogs("netbox_interface_name_rules"): + with connection.execute_wrapper(reject_interface_updates), self.assertLogs("netbox_interface_name_rules"): for callback in callbacks: callback() - (entry,) = _journal(card) + (entry,) = journal(card) self.assertEqual(entry.kind, JournalEntryKindChoices.KIND_DANGER) self.assertIn(f"`et-1/0/1` to `et-1/2/1`: {TAKEN}", entry.comments) self.assertIn("injected reapply failure", entry.comments) self.assertEqual((self._names(blocked), self._names(failed)), (["et-1/0/1"], ["et-1/0/2"])) - self.assertEqual((_journal(blocked), _journal(failed)), ([], [])) + self.assertEqual((journal(blocked), journal(failed)), ([], [])) def test_a_type_change_and_a_bay_edit_rename_the_nested_modules_from_the_naming_before_the_edit(self): bay = self._bay(self.device) @@ -283,7 +251,7 @@ def test_an_outer_edit_undone_around_a_nested_edit_renames_the_nested_module_for self._save_edit(bay, position="0") self.assertEqual(self._names(optic), ["et-1/0/3"]) - self.assertEqual((_journal(card), _journal(optic)), ([], [])) + self.assertEqual((journal(card), journal(optic)), ([], [])) def test_a_nested_edit_before_an_outer_edit_renames_the_nested_module_without_a_report(self): bay, card, port, optic = self._card_with_optic() @@ -293,7 +261,7 @@ def test_a_nested_edit_before_an_outer_edit_renames_the_nested_module_without_a_ self._save_edit(bay, position="2") self.assertEqual(self._names(optic), ["et-1/2/3"]) - self.assertEqual((_journal(card), _journal(optic)), ([], [])) + self.assertEqual((journal(card), journal(optic)), ([], [])) def test_an_outer_reapply_that_runs_first_renames_the_nested_module_from_its_earliest_naming(self): bay, card, port, optic = self._card_with_optic() @@ -306,11 +274,11 @@ def test_an_outer_reapply_that_runs_first_renames_the_nested_module_from_its_ear self._save_edit(bay, position="2") self.assertEqual(self._names(optic), ["et-1/2/3"]) - self.assertEqual((_journal(card), _journal(optic)), ([], [])) + self.assertEqual((journal(card), journal(optic)), ([], [])) def test_a_card_and_an_optic_installed_before_an_edit_of_the_card_bay_build_the_optic_family(self): flat_optic_type = self._module_type("Flat Optic", "{module}") - _flat_rule(flat_optic_type, "x-{slot}/{bay_position}:{channel}") + flat_rule(flat_optic_type, "x-{slot}/{bay_position}:{channel}") bay = self._bay(self.device) with self.captureOnCommitCallbacks(execute=True), transaction.atomic(): @@ -320,11 +288,11 @@ def test_a_card_and_an_optic_installed_before_an_edit_of_the_card_bay_build_the_ self._save_edit(bay, position="2") self.assertEqual(self._names(optic), ["x-2/1:0", "x-2/1:1"]) - self.assertEqual((_journal(card), _journal(optic)), ([], [])) + self.assertEqual((journal(card), journal(optic)), ([], [])) def test_an_optic_installed_in_a_card_before_an_edit_of_the_card_bay_builds_its_family_once(self): flat_optic_type = self._module_type("Flat Optic", "{module}") - _flat_rule(flat_optic_type, "x-{slot}/{bay_position}:{channel}") + flat_rule(flat_optic_type, "x-{slot}/{bay_position}:{channel}") bay = self._bay(self.device) card, port = self._install_card(self._card_type("Card", "1"), bay) @@ -333,11 +301,11 @@ def test_an_optic_installed_in_a_card_before_an_edit_of_the_card_bay_builds_its_ self._save_edit(bay, position="2") self.assertEqual(self._names(optic), ["x-2/1:0", "x-2/1:1"]) - self.assertEqual((_journal(card), _journal(optic)), ([], [])) + self.assertEqual((journal(card), journal(optic)), ([], [])) def test_an_optic_installed_before_edits_of_the_card_bay_and_its_own_bay_builds_its_family_once(self): flat_optic_type = self._module_type("Flat Optic", "{module}") - _flat_rule(flat_optic_type, "x-{slot}/{bay_position}:{channel}") + flat_rule(flat_optic_type, "x-{slot}/{bay_position}:{channel}") bay = self._bay(self.device) card, port = self._install_card(self._card_type("Card", "1"), bay) @@ -347,13 +315,13 @@ def test_an_optic_installed_before_edits_of_the_card_bay_and_its_own_bay_builds_ self._save_edit(port, position="3") self.assertEqual(self._names(optic), ["x-2/3:0", "x-2/3:1"]) - self.assertEqual((_journal(card), _journal(optic)), ([], [])) + self.assertEqual((journal(card), journal(optic)), ([], [])) def test_an_optic_moved_into_a_card_before_an_edit_of_the_card_bay_builds_its_family_once(self): card_type = self._card_type("Card", "1") optic_type = self._module_type("Scoped Optic", "{module}") InterfaceNameRule.objects.create(module_type=optic_type, name_template="p{bay_position}") - _flat_rule(optic_type, "x-{slot}/{bay_position}:{channel}", parent_module_type=card_type) + flat_rule(optic_type, "x-{slot}/{bay_position}:{channel}", parent_module_type=card_type) bay = self._bay(self.device) card, port = self._install_card(card_type, bay) optic = self._install(optic_type, self._bay(self.device, "Bay 1")) @@ -364,7 +332,7 @@ def test_an_optic_moved_into_a_card_before_an_edit_of_the_card_bay_builds_its_fa self._save_edit(bay, position="2") self.assertEqual(self._names(optic), ["x-2/1:0", "x-2/1:1"]) - self.assertEqual((_journal(card), _journal(optic)), ([], [])) + self.assertEqual((journal(card), journal(optic)), ([], [])) def test_an_optic_whose_raw_name_reads_the_card_bay_is_renamed_after_an_edit_of_that_bay(self): chained_optic_type = self._module_type("Chained Optic", "{module}/{module}") @@ -379,7 +347,7 @@ def test_an_optic_whose_raw_name_reads_the_card_bay_is_renamed_after_an_edit_of_ self._save_edit(bay, position="2") self.assertEqual(self._names(optic), ["et-1/2/1"]) - self.assertEqual((_journal(card), _journal(optic)), ([], [])) + self.assertEqual((journal(card), journal(optic)), ([], [])) @skipUnless(supports_module_moves(), REQUIRES_SUBTREE_MOVES) def test_the_bay_post_saves_netbox_sends_in_a_move_are_not_bay_triggers(self): @@ -387,7 +355,7 @@ def test_the_bay_post_saves_netbox_sends_in_a_move_are_not_bay_triggers(self): optic = self._install(self.optic_type, port) self.assertEqual(self._names(optic), ["et-1/0/0"]) - with _naming_reads() as reads, self.captureOnCommitCallbacks() as callbacks: + with naming_reads() as reads, self.captureOnCommitCallbacks() as callbacks: self._save_move(card, self._bay(self.device, "Bay 2")) (runner,) = [callback for callback in callbacks if isinstance(callback, PlanRunner)] for callback in callbacks: @@ -427,14 +395,14 @@ def test_a_nested_module_with_its_own_trigger_reports_on_the_module_in_the_outer bay, card, port, optic = self._card_with_optic() Interface.objects.create(device=self.device, name="et-1/2/3", type=PLAIN_TYPE) - with _module_reapplies() as reapplies, self.captureOnCommitCallbacks(execute=True), transaction.atomic(): + with module_reapplies() as reapplies, self.captureOnCommitCallbacks(execute=True), transaction.atomic(): self._save_edit(bay, position="2") self._save_edit(port, position="3") self.assertEqual((self._names(optic), reapplies.call_count), (["et-1/0/1"], 2)) - (entry,) = _journal(card) + (entry,) = journal(card) self.assertIn(f"`et-1/0/1` to `et-1/2/3`: {TAKEN}", entry.comments) - self.assertEqual(_journal(optic), []) + self.assertEqual(journal(optic), []) def test_a_nested_type_change_and_an_outer_edit_reapply_the_nested_module_once_as_a_type_change(self): bay, card, _port, optic = self._card_with_optic() @@ -444,20 +412,20 @@ def test_a_nested_type_change_and_an_outer_edit_reapply_the_nested_module_once_a ) Interface.objects.create(device=self.device, name="ge-1/2/1", type=PLAIN_TYPE) - with _module_reapplies() as reapplies, self.captureOnCommitCallbacks(execute=True), transaction.atomic(): + with module_reapplies() as reapplies, self.captureOnCommitCallbacks(execute=True), transaction.atomic(): optic.module_type = other_optic_type optic.save() self._save_edit(bay, position="2") self.assertEqual((self._names(optic), reapplies.call_count), (["et-1/0/1"], 2)) - (entry,) = _journal(card) + (entry,) = journal(card) self.assertIn(f"`et-1/0/1` to `ge-1/2/1`: {TAKEN}", entry.comments) - self.assertEqual(_journal(optic), []) + self.assertEqual(journal(optic), []) def test_an_outer_edit_and_a_move_of_the_nested_module_out_reapply_it_once(self): bay, _card, _port, optic = self._card_with_optic() - with _module_reapplies() as reapplies, self.captureOnCommitCallbacks(execute=True), transaction.atomic(): + with module_reapplies() as reapplies, self.captureOnCommitCallbacks(execute=True), transaction.atomic(): self._save_edit(bay, position="2") self._save_move(optic, self._bay(self.device, "Bay 1")) @@ -467,25 +435,25 @@ def test_an_outer_edit_and_a_move_of_the_nested_module_out_reapply_it_once(self) def test_an_outer_move_and_a_nested_bay_edit_reapply_each_module_once(self): _bay, card, port, optic = self._card_with_optic() - with _module_reapplies() as reapplies, self.captureOnCommitCallbacks(execute=True), transaction.atomic(): + with module_reapplies() as reapplies, self.captureOnCommitCallbacks(execute=True), transaction.atomic(): self._save_move(card, self._bay(self.device, "Bay 2")) port.refresh_from_db() self._save_edit(port, position="3") self.assertEqual((self._names(optic), reapplies.call_count), (["et-1/2/3"], 2)) - self.assertEqual((_journal(card), _journal(optic)), ([], [])) + self.assertEqual((journal(card), journal(optic)), ([], [])) @skipUnless(supports_module_moves(), REQUIRES_SUBTREE_MOVES) def test_a_nested_bay_edit_and_an_outer_move_reapply_each_module_once(self): _bay, card, port, optic = self._card_with_optic() - with _module_reapplies() as reapplies, self.captureOnCommitCallbacks(execute=True), transaction.atomic(): + with module_reapplies() as reapplies, self.captureOnCommitCallbacks(execute=True), transaction.atomic(): self._save_edit(port, position="3") self._save_move(card, self._bay(self.device, "Bay 2")) port.refresh_from_db() self.assertEqual((self._names(optic), reapplies.call_count), ([f"et-1/2/{port.position}"], 2)) - self.assertEqual((_journal(card), _journal(optic)), ([], [])) + self.assertEqual((journal(card), journal(optic)), ([], [])) def test_an_optic_installed_while_the_card_bay_is_edited_and_restored_is_named_from_its_raw_name(self): chained_optic_type = self._module_type("Chained Optic", "{module}/{module}") @@ -501,7 +469,7 @@ def test_an_optic_installed_while_the_card_bay_is_edited_and_restored_is_named_f self._save_edit(bay, position="0") self.assertEqual(self._names(optic), ["et-1/0/1"]) - self.assertEqual((_journal(card), _journal(optic)), ([], [])) + self.assertEqual((journal(card), journal(optic)), ([], [])) def test_a_rolled_back_nested_edit_leaves_the_outer_edit_to_rename_the_nested_module(self): bay, card, port, optic = self._card_with_optic() @@ -513,7 +481,7 @@ def test_a_rolled_back_nested_edit_leaves_the_outer_edit_to_rename_the_nested_mo raise RuntimeError("roll back the savepoint") self.assertEqual(self._names(optic), ["et-1/2/1"]) - self.assertEqual((_journal(card), _journal(optic)), ([], [])) + self.assertEqual((journal(card), journal(optic)), ([], [])) def test_a_rolled_back_outer_edit_leaves_the_nested_edit_to_rename_its_module(self): bay, card, port, optic = self._card_with_optic() @@ -525,7 +493,7 @@ def test_a_rolled_back_outer_edit_leaves_the_nested_edit_to_rename_its_module(se raise RuntimeError("roll back the savepoint") self.assertEqual(self._names(optic), ["et-1/0/3"]) - self.assertEqual((_journal(card), _journal(optic)), ([], [])) + self.assertEqual((journal(card), journal(optic)), ([], [])) def test_a_naming_read_in_a_rolled_back_savepoint_is_not_used(self): other_card_type = self._card_type("Other Card", "1") @@ -554,7 +522,7 @@ def test_a_naming_read_in_a_rolled_back_savepoint_is_not_used(self): self._save_edit(bay, position="3") self.assertEqual((self._names(optic), self._names(other)), (["a-3/1"], ["et-1/0/1"])) - self.assertEqual((_journal(card), _journal(optic)), ([], [])) + self.assertEqual((journal(card), journal(optic)), ([], [])) class BayEditTransactionTest(BayEditTestCase): @@ -564,7 +532,7 @@ def test_several_edits_of_one_bay_reapply_once_from_its_earliest_state(self): bay = self._bay(self.device) module = self._install(self.plain_type, bay) - with _module_reapplies() as reapplies, self.captureOnCommitCallbacks(execute=True), transaction.atomic(): + with module_reapplies() as reapplies, self.captureOnCommitCallbacks(execute=True), transaction.atomic(): self._save_edit(bay, position="5") self._save_edit(bay, position="7") @@ -578,20 +546,20 @@ def test_an_edit_the_same_transaction_undoes_reapplies_nothing(self): edited_bay = self._bay(self.device, "Bay 1") edited = self._install(self.plain_type, edited_bay) - with _module_reapplies() as reapplies, self.captureOnCommitCallbacks(execute=True), transaction.atomic(): + with module_reapplies() as reapplies, self.captureOnCommitCallbacks(execute=True), transaction.atomic(): self._save_edit(returned_bay, position="5") self._save_edit(returned_bay, position="0") self._save_edit(edited_bay, position="4") self.assertEqual(reapplies.call_count, 1) self.assertEqual((self._names(returned), self._names(edited)), (["operator-name"], ["et-1/0/4"])) - self.assertEqual(_journal(returned), []) + self.assertEqual(journal(returned), []) def test_an_edit_in_a_rolled_back_savepoint_causes_no_reapply_and_a_later_edit_reapplies_once(self): bay = self._bay(self.device) module = self._install(self.plain_type, bay) - with _module_reapplies() as reapplies, self.captureOnCommitCallbacks(execute=True), transaction.atomic(): + with module_reapplies() as reapplies, self.captureOnCommitCallbacks(execute=True), transaction.atomic(): with self.assertRaises(RuntimeError), transaction.atomic(): self._save_edit(bay, position="9") raise RuntimeError("roll back the savepoint") @@ -609,7 +577,7 @@ def test_an_edit_in_a_rolled_back_savepoint_schedules_no_reapply(self): with self.assertRaises(RuntimeError), transaction.atomic(): self._save_edit(bay, position="9") raise RuntimeError("roll back the savepoint") - with _module_reapplies() as reapplies: + with module_reapplies() as reapplies: for callback in callbacks: callback() @@ -620,7 +588,7 @@ def test_an_edit_rolled_back_after_an_earlier_edit_keeps_the_earlier_reapply(sel bay = self._bay(self.device) module = self._install(self.plain_type, bay) - with _module_reapplies() as reapplies, self.captureOnCommitCallbacks(execute=True), transaction.atomic(): + with module_reapplies() as reapplies, self.captureOnCommitCallbacks(execute=True), transaction.atomic(): self._save_edit(bay, position="5") with self.assertRaises(RuntimeError), transaction.atomic(): self._save_edit(bay, position="7") @@ -654,7 +622,7 @@ def test_a_bay_edit_joins_a_plan_that_a_save_outside_a_capture_left_pending(self def test_an_install_and_an_edit_of_its_bay_name_the_module_for_the_edited_bay(self): bay = self._bay(self.device) - with _module_reapplies() as reapplies, self.captureOnCommitCallbacks(execute=True), transaction.atomic(): + with module_reapplies() as reapplies, self.captureOnCommitCallbacks(execute=True), transaction.atomic(): module = Module.objects.create(device=self.device, module_bay=bay, module_type=self.plain_type) self._save_edit(bay, position="5") @@ -663,7 +631,7 @@ def test_an_install_and_an_edit_of_its_bay_name_the_module_for_the_edited_bay(se def test_an_install_and_an_edit_of_its_bay_under_a_flat_rule_build_the_family(self): flat_type = self._module_type("Flat", "{module}") - _flat_rule(flat_type, "f-{bay_position}:{channel}") + flat_rule(flat_type, "f-{bay_position}:{channel}") bay = self._bay(self.device) with self.captureOnCommitCallbacks(execute=True), transaction.atomic(): @@ -671,13 +639,13 @@ def test_an_install_and_an_edit_of_its_bay_under_a_flat_rule_build_the_family(se self._save_edit(bay, position="5") self.assertEqual(self._names(module), ["f-5:0", "f-5:1"]) - self.assertEqual(_journal(module), []) + self.assertEqual(journal(module), []) def test_a_module_moved_into_a_bay_that_is_then_edited_reapplies_once_for_the_edited_bay(self): module = self._install(self.plain_type, self._bay(self.device)) bay = self._bay(self.device, "Bay 1") - with _module_reapplies() as reapplies, self.captureOnCommitCallbacks(execute=True), transaction.atomic(): + with module_reapplies() as reapplies, self.captureOnCommitCallbacks(execute=True), transaction.atomic(): self._save_move(module, bay) self._save_edit(bay, position="4") @@ -757,8 +725,8 @@ def _assert_reapplied_once(self, model, chassis_first, before, after): self.assertEqual( (self._names(optic), self._names(other)), ([after.format(raw=self._raw_name(optic))], ["et-3/0/10"]) ) - self.assertEqual(_reapplied(reapplies), sorted((card.pk, optic.pk, other.pk))) - self.assertEqual((_journal(card), _journal(optic), _journal(self.device)), ([], [], [])) + self.assertEqual(reapplied(reapplies), sorted((card.pk, optic.pk, other.pk))) + self.assertEqual((journal(card), journal(optic), journal(self.device)), ([], [], [])) def test_a_bay_edit_then_a_chassis_position_change_reapply_a_plain_rule_once(self): self._assert_reapplied_once("Plain", False, "et-1/0/1", "et-3/2/1") @@ -784,10 +752,10 @@ def _assert_a_collision_is_reported_once(self, chassis_first): reapplies = self._edit_with(bay, self._change_the_chassis_position, chassis_first) - self.assertEqual((self._names(optic), _reapplied(reapplies).count(optic.pk)), (["et-1/0/1"], 1)) - (entry,) = _journal(card) + self.assertEqual((self._names(optic), reapplied(reapplies).count(optic.pk)), (["et-1/0/1"], 1)) + (entry,) = journal(card) self.assertEqual(entry.comments.count(f"`et-1/0/1` to `et-3/2/1`: {TAKEN}"), 1) - self.assertEqual((_journal(optic), _journal(self.device)), ([], [])) + self.assertEqual((journal(optic), journal(self.device)), ([], [])) def test_a_bay_edit_then_a_chassis_position_change_report_a_collision_once(self): self._assert_a_collision_is_reported_once(chassis_first=False) @@ -801,9 +769,9 @@ def _assert_leaving_renames_nothing_and_reports_once(self, leave_first): reapplies = self._edit_with(bay, self._leave_the_chassis, leave_first) self.assertEqual((self._names(optic), self._names(other)), (["x10/1"], ["et-1/0/10"])) - self.assertEqual(_reapplied(reapplies), sorted((card.pk, optic.pk, other.pk))) - (card_entry,) = _journal(card) - (device_entry,) = _journal(self.device) + self.assertEqual(reapplied(reapplies), sorted((card.pk, optic.pk, other.pk))) + (card_entry,) = journal(card) + (device_entry,) = journal(self.device) self.assertEqual(card_entry.comments.count(f"`x10/1`: {UNAVAILABLE}"), 1) self.assertIn(f"`et-1/0/10`: {UNAVAILABLE}", device_entry.comments) self.assertNotIn("x10/1", device_entry.comments) @@ -825,7 +793,7 @@ def _assert_joining_renames_once(self, join_first): reapplies = self._edit_with(bay, functools.partial(self._join_the_chassis, chassis), join_first) self.assertEqual((self._names(optic), self._names(other)), (["x32/1"], ["et-3/0/10"])) - self.assertEqual(_reapplied(reapplies), sorted((card.pk, optic.pk, other.pk))) + self.assertEqual(reapplied(reapplies), sorted((card.pk, optic.pk, other.pk))) self.assertEqual(JournalEntry.objects.count(), entries) def test_a_bay_edit_then_joining_a_chassis_rename_each_module_once(self): @@ -863,8 +831,8 @@ def _install_after_a_device_change(self, device_change, module_type): def test_a_chassis_position_change_an_install_and_a_bay_edit_recognise_the_raw_name_of_the_install(self): module, reapplies = self._install_after_a_device_change(self._change_the_chassis_position, self.token_type) - self.assertEqual((self._names(module), _reapplied(reapplies)), (["et-3/2"], [module.pk])) - self.assertEqual((_journal(module), _journal(self.device)), ([], [])) + self.assertEqual((self._names(module), reapplied(reapplies)), (["et-3/2"], [module.pk])) + self.assertEqual((journal(module), journal(self.device)), ([], [])) @skipUnless(supports_vc_position_token(), REQUIRES_VC_POSITION_TOKEN) def test_an_install_between_two_chassis_position_changes_is_recognised_at_the_position_of_the_install(self): @@ -878,8 +846,8 @@ def test_an_install_between_two_chassis_position_changes_is_recognised_at_the_po ) (module,) = installed - self.assertEqual((self._names(module), _reapplied(reapplies)), (["et-5/2"], [module.pk])) - self.assertEqual((_journal(module), _journal(self.device)), ([], [])) + self.assertEqual((self._names(module), reapplied(reapplies)), (["et-5/2"], [module.pk])) + self.assertEqual((journal(module), journal(self.device)), ([], [])) @skipUnless(supports_vc_position_token(), REQUIRES_VC_POSITION_TOKEN) def test_joining_a_chassis_an_install_and_a_bay_edit_recognise_the_raw_name_of_the_install(self): @@ -891,21 +859,21 @@ def test_joining_a_chassis_an_install_and_a_bay_edit_recognise_the_raw_name_of_t functools.partial(self._join_the_chassis, chassis), self.token_type ) - self.assertEqual((self._names(module), _reapplied(reapplies)), (["et-3/2"], [module.pk])) - self.assertEqual((_journal(module), _journal(self.device)), ([], [])) + self.assertEqual((self._names(module), reapplied(reapplies)), (["et-3/2"], [module.pk])) + self.assertEqual((journal(module), journal(self.device)), ([], [])) @skipUnless(supports_vc_position_token(), REQUIRES_VC_POSITION_TOKEN) def test_leaving_a_chassis_an_install_and_a_bay_edit_recognise_the_raw_name_of_the_install(self): module, reapplies = self._install_after_a_device_change(self._leave_the_chassis, self.token_base_type) - self.assertEqual((self._names(module), _reapplied(reapplies)), (["p0/2"], [module.pk])) - self.assertEqual((_journal(module), _journal(self.device)), ([], [])) + self.assertEqual((self._names(module), reapplied(reapplies)), (["p0/2"], [module.pk])) + self.assertEqual((journal(module), journal(self.device)), ([], [])) def _assert_an_installed_raw_name_is_renamed_once(self, chassis_first): module, reapplies = self._install_and_edit_a_token_module(self._change_the_chassis_position, chassis_first) - self.assertEqual((self._names(module), _reapplied(reapplies)), (["et-3/2"], [module.pk])) - self.assertEqual((_journal(module), _journal(self.device)), ([], [])) + self.assertEqual((self._names(module), reapplied(reapplies)), (["et-3/2"], [module.pk])) + self.assertEqual((journal(module), journal(self.device)), ([], [])) @skipUnless(supports_vc_position_token(), REQUIRES_VC_POSITION_TOKEN) def test_an_install_a_bay_edit_and_a_chassis_position_change_rename_a_raw_name_that_reads_the_position(self): @@ -926,7 +894,7 @@ def _assert_a_raw_name_off_the_chassis_is_renamed_after_a_join(self, join_first) reapplies = self._edit_with(bay, functools.partial(self._join_the_chassis, chassis), join_first) - self.assertEqual((self._names(module), _reapplied(reapplies)), (["et-3/2"], [module.pk])) + self.assertEqual((self._names(module), reapplied(reapplies)), (["et-3/2"], [module.pk])) self.assertEqual(JournalEntry.objects.count(), entries) @skipUnless(supports_vc_position_token(), REQUIRES_VC_POSITION_TOKEN) @@ -946,7 +914,7 @@ def test_a_bay_save_fails_with_the_previous_state_read_error(self): self._install(self.plain_type, bay) with ( - _previous_state_read_fails(BAY_STATE_READ) as replaced, + previous_state_read_fails(BAY_STATE_READ) as replaced, self.assertRaisesMessage(DataError, "division by zero"), transaction.atomic(), ): @@ -960,7 +928,7 @@ def test_a_naming_read_that_fails_fails_the_bay_edit_with_its_error(self): module = self._install(self.plain_type, bay) with ( - connection.execute_wrapper(_fail_the_naming_read), + connection.execute_wrapper(fail_the_naming_read), self.assertRaisesMessage(DataError, "division by zero"), transaction.atomic(), ): @@ -970,7 +938,7 @@ def test_a_naming_read_that_fails_fails_the_bay_edit_with_its_error(self): self.assertEqual(self._names(module), ["et-1/0/0"]) -class BayEditAPITest(_MoveFixture, APITestCase): +class BayEditAPITest(MoveFixture, APITestCase): """A REST API bay edit reaches the rename trigger through NetBox's own write path.""" model = ModuleBay @@ -993,4 +961,4 @@ def test_patching_the_bay_position_renames_the_installed_module(self): self.assertEqual(response.status_code, status.HTTP_200_OK, response.data) self.assertEqual(self._names(module), ["et-1/0/5"]) - self.assertEqual(_journal(module), []) + self.assertEqual(journal(module), []) diff --git a/netbox_interface_name_rules/tests/test_branch_jobs.py b/netbox_interface_name_rules/tests/test_branch_jobs.py new file mode 100644 index 00000000..36ff6e50 --- /dev/null +++ b/netbox_interface_name_rules/tests/test_branch_jobs.py @@ -0,0 +1,201 @@ +# SPDX-License-Identifier: Apache-2.0 +# Copyright (C) 2025 Marcin Zieba +"""Plugin jobs enqueued in a real netbox-branching branch run in that branch, and only there. + +Each test enqueues its job through the Apply page with netbox-branching's cookie, then runs the job +from what the queue stored, as a worker does. A background REST request for a rule in a branch is +refused before it writes. +""" + +from core.choices import JobStatusChoices +from core.models import Job, ObjectChange +from django.conf import settings +from django.urls import reverse + +from netbox_interface_name_rules.api.views import BACKGROUND_IN_A_BRANCH +from netbox_interface_name_rules.models import InterfaceNameRule +from netbox_interface_name_rules.tests.branch_cases import FLAT_NAMES, BranchWriteCase, ConversionCase, PlainModuleCase +from netbox_interface_name_rules.tests.helpers import ( + branch_cookie, + make_manufacturer, + make_module_type, + names_of, + queued_job, + register_a_worker, +) + +COMPLETED = JobStatusChoices.STATUS_COMPLETED +ERRORED = JobStatusChoices.STATUS_ERRORED + + +class _JobCase(BranchWriteCase): + """Enqueue a job through the Apply page in the branch, and run it from the queue as a worker does.""" + + def enqueue(self, action): + """Post *action* on the Apply page of the rule and return the one job that it enqueued.""" + before = set(Job.objects.values_list("pk", flat=True)) + + response = self.client.post(self.apply_url(self.rule), {"action": action}) + + self.assertEqual(response.status_code, 302) + [job] = Job.objects.exclude(pk__in=before) + return job + + def run_queued(self, job): + """Run *job* from the queue's record as a worker does, outside the branch, and return it reloaded.""" + queued_job(self, job).perform() + job.refresh_from_db() + return job + + def make_not_ready(self): + """Mark the branch merged, as a merge does, so netbox-branching no longer activates it.""" + from netbox_branching.choices import BranchStatusChoices + from netbox_branching.models import Branch + + Branch.objects.filter(pk=self.branch.pk).update(status=BranchStatusChoices.MERGED) + + def alias_error(self): + """Return the error of a job that started on main although it was enqueued in the branch.""" + return repr(RuntimeError(f"The write alias is 'default', but the operation expects {self.alias!r}.")) + + +class ApplyJobInABranchTest(_JobCase, PlainModuleCase): + PREFIX = "BrJobApply" + + def build(self): + super().build() + self.rule = InterfaceNameRule.objects.create( + module_type=self.module_type, name_template="et-0/0/{bay_position}" + ) + + def test_the_job_stores_the_branch_and_its_alias_and_no_session_cookie(self): + job = self.enqueue("background") + + queued = queued_job(self, job) + + self.assertEqual( + queued.kwargs, + { + "job": job, + "rule_id": self.rule.pk, + "branch_schema_id": self.branch.schema_id, + "expected_alias": self.alias, + }, + ) + session = self.client.cookies[settings.SESSION_COOKIE_NAME].value + self.assertNotIn(session.encode(), queued.data) + + def test_the_job_renames_in_the_branch_only(self): + job = self.run_queued(self.enqueue("background")) + + self.assertEqual((job.status, job.error), (COMPLETED, "")) + with self.in_branch(): + self.assertEqual(names_of(self.module), ["et-0/0/0"]) + self.assertEqual(names_of(self.module), ["0"]) + self.assertEqual(self.branch_updates_of(self.interface), [("0", "et-0/0/0")]) + self.assertEqual(list(self.change_diffs().values_list("object_id", flat=True)), [self.interface.pk]) + + def test_a_job_whose_branch_is_no_longer_ready_fails_naming_both_aliases_and_renames_nothing(self): + job = self.enqueue("background") + self.make_not_ready() + + job = self.run_queued(job) + + self.assertEqual((job.status, job.error), (ERRORED, self.alias_error())) + with self.in_branch(): + self.assertEqual(names_of(self.module), ["0"]) + self.assertEqual(names_of(self.module), ["0"]) + self.assertFalse(self.change_diffs().exists()) + + +class StartCheckBeforeTheRuleReadTest(_JobCase, PlainModuleCase): + """The rule exists in the branch only, so a job that read it on main would find no rule and complete.""" + + PREFIX = "BrJobRuleRead" + + def setUp(self): + super().setUp() + with self.in_branch(): + self.rule = InterfaceNameRule.objects.create( + module_type=self.module_type, name_template="et-0/0/{bay_position}" + ) + + def test_a_job_whose_branch_is_no_longer_ready_fails_before_it_reads_the_rule(self): + job = self.enqueue("background") + self.make_not_ready() + + job = self.run_queued(job) + + self.assertEqual((job.status, job.error), (ERRORED, self.alias_error())) + with self.in_branch(): + self.assertEqual(names_of(self.module), ["0"]) + + +class ConvertJobInABranchTest(_JobCase, ConversionCase): + PREFIX = "BrJobConvert" + + def test_the_job_converts_in_the_branch_only(self): + job = self.run_queued(self.enqueue("convert_background")) + + self.assertEqual((job.status, job.error), (COMPLETED, "")) + with self.in_branch(): + self.assertEqual(names_of(self.module), ["et-0/0/3", *FLAT_NAMES]) + self.assertEqual(names_of(self.module), list(FLAT_NAMES)) + + def test_a_job_whose_branch_is_no_longer_ready_fails_and_converts_nothing(self): + job = self.enqueue("convert_background") + self.make_not_ready() + + job = self.run_queued(job) + + self.assertEqual((job.status, job.error), (ERRORED, self.alias_error())) + with self.in_branch(): + self.assertEqual(names_of(self.module), list(FLAT_NAMES)) + self.assertEqual(names_of(self.module), list(FLAT_NAMES)) + + +class BackgroundRuleRequestTest(BranchWriteCase): + """NetBox runs a background REST request on main, so the rule endpoints refuse one in a branch.""" + + PREFIX = "BrBackground" + TEMPLATE = "et-0/0/{bay_position}" + EDITED = "xe-0/0/{bay_position}" + + def build(self): + self.module_type = make_module_type(make_manufacturer(self.PREFIX), self.PREFIX) + self.rule = InterfaceNameRule.objects.create(module_type=self.module_type, name_template=self.TEMPLATE) + + def setUp(self): + super().setUp() + del self.client.cookies[branch_cookie()] + # Without the refusal, NetBox would accept the request, because a worker is registered. + register_a_worker(self) + + def background(self, method, payload, **headers): + url = reverse("plugins-api:netbox_interface_name_rules-api:interfacenamerule-list") + return getattr(self.client, method)( + f"{url}?background=true", payload, content_type="application/json", headers=headers + ) + + def rules(self): + return list(InterfaceNameRule.objects.values_list("pk", "name_template")) + + def test_a_background_rule_request_in_a_branch_is_refused_and_writes_nothing_on_either_alias(self): + requests = { + "post": [{"module_type": self.module_type.pk, "name_template": self.EDITED}], + "put": [{"id": self.rule.pk, "module_type": self.module_type.pk, "name_template": self.EDITED}], + "patch": [{"id": self.rule.pk, "name_template": self.EDITED}], + "delete": [{"id": self.rule.pk}], + } + for method, payload in requests.items(): + with self.subTest(method=method): + response = self.background(method, payload, **{"X-NetBox-Branch": self.branch.schema_id}) + + self.assertEqual((response.status_code, response.json()), (400, [BACKGROUND_IN_A_BRANCH])) + + self.assertFalse(Job.objects.exists()) + self.assertEqual(self.rules(), [(self.rule.pk, self.TEMPLATE)]) + with self.in_branch(): + self.assertEqual(self.rules(), [(self.rule.pk, self.TEMPLATE)]) + self.assertFalse(ObjectChange.objects.using(self.alias).exists()) + self.assertFalse(self.change_diffs().exists()) diff --git a/netbox_interface_name_rules/tests/test_branch_replay.py b/netbox_interface_name_rules/tests/test_branch_replay.py new file mode 100644 index 00000000..19205932 --- /dev/null +++ b/netbox_interface_name_rules/tests/test_branch_replay.py @@ -0,0 +1,378 @@ +# SPDX-License-Identifier: Apache-2.0 +# Copyright (C) 2025 Marcin Zieba +"""A netbox-branching merge, revert or sync replays logged changes, and the rename triggers do nothing then. + +A test starts the operation on netbox-branching's page of the branch, as an operator does, and runs +the job that the page enqueued from the queue, as a worker does, or it calls the operation as a shell +does. netbox-branching replays only the changes that NetBox logged, so the changes that a test +replays go through requests. +""" + +from core.choices import JobStatusChoices +from core.models import Job, ObjectChange +from dcim.models import Device, Interface, InterfaceTemplate, Module, VirtualChassis +from django.db import transaction +from django.urls import reverse +from extras.models import JournalEntry + +from netbox_interface_name_rules.models import InterfaceNameRule +from netbox_interface_name_rules.tests.branch_cases import BranchWriteCase, KeptChannelCase +from netbox_interface_name_rules.tests.helpers import ( + PLAIN_TYPE, + branch_cookie, + install_form, + make_device, + make_device_type, + make_manufacturer, + make_module_bay_templates, + make_module_type, + queued_job, +) + +COMPLETED = JobStatusChoices.STATUS_COMPLETED +ERRORED = JobStatusChoices.STATUS_ERRORED +ITERATIVE = "iterative" +SQUASH = "squash" +POSITIONS = (0, 1, 2) +# The names of each bay after the branch installed a module in bay 0 and moved the module of bay 1 to bay 2. +BRANCH_NAMES = {0: ["a0", "b0"], 1: [], 2: ["a2", "b2"]} +NAMES_BEFORE = {0: [], 1: ["a1", "b1"], 2: []} +# The names that NetBox gives the channelized family in bay 1. +RAW_FAMILY = ["1", "1:1", "1:2", "1:3", "1:4"] + + +def merging(strategy, commit=True): + """Return the form that merges the branch with *strategy*, or only tries it when *commit* is false.""" + return {"merge_strategy": strategy, **({"commit": "on"} if commit else {})} + + +class _ReplayCase(BranchWriteCase): + """Run netbox-branching's merge, revert and sync as an operator and a worker do.""" + + def enqueue(self, action, **form): + """Post *form* on netbox-branching's *action* page of the branch and return the job that it enqueued.""" + before = set(Job.objects.values_list("pk", flat=True)) + + response = self.client.post( + reverse(f"plugins:netbox_branching:branch_{action}", kwargs={"pk": self.branch.pk}), form + ) + + self.assertEqual(response.status_code, 302) + [job] = Job.objects.exclude(pk__in=before) + return job + + def run_queued(self, job): + """Run *job* from the queue's record as a worker does, and return it reloaded.""" + queued_job(self, job).perform() + job.refresh_from_db() + return job + + def act(self, action, **form): + """Post *form* on netbox-branching's *action* page of the branch, run the job it enqueued, and return it.""" + return self.run_queued(self.enqueue(action, **form)) + + def on_main(self): + """Send the next requests of the test client to main.""" + del self.client.cookies[branch_cookie()] + + def last_change_on_main(self): + """Return the primary key of the newest change that NetBox logged on main.""" + return ObjectChange.objects.order_by("-pk").values_list("pk", flat=True).first() or 0 + + def assert_only_replayed_changes_on_main(self, since): + """Assert that each change logged on main after *since* replays a change of the branch, and no journal entry.""" + replayed = set(ObjectChange.objects.using(self.alias).values_list("request_id", flat=True)) + logged = set(ObjectChange.objects.filter(pk__gt=since).values_list("request_id", flat=True)) + + self.assertTrue(logged) + self.assertLessEqual(logged, replayed) + self.assertFalse(JournalEntry.objects.exists()) + + +class _InstallAndMoveCase(_ReplayCase): + """A device with three module bays, and a module in bay 1 whose interfaces NetBox named ``a1`` and ``b1``.""" + + def build(self): + manufacturer = make_manufacturer(self.PREFIX) + device_type = make_device_type(manufacturer, self.PREFIX) + make_module_bay_templates(device_type, tuple(f"Bay {position}" for position in POSITIONS)) + self.device = make_device(self.PREFIX, device_type) + self.module_type = make_module_type(manufacturer, self.PREFIX) + for template in ("a{module}", "b{module}"): + InterfaceTemplate.objects.create(module_type=self.module_type, name=template, type=PLAIN_TYPE) + with transaction.atomic(): + self.module = Module.objects.create( + device=self.device, module_bay=self.bay(1), module_type=self.module_type + ) + + def add_rule(self): + return InterfaceNameRule.objects.create(module_type=self.module_type, name_template="{base}.br") + + def install_and_move(self): + """Install a module in bay 0 through the UI and move the module of bay 1 to bay 2 through the REST API.""" + response = self.client.post(reverse("dcim:module_add"), install_form(self.bay(0), self.module_type)) + self.assertEqual(response.status_code, 302) + response = self.client.patch( + reverse("dcim-api:module-detail", kwargs={"pk": self.module.pk}), + {"module_bay": self.bay(2).pk}, + content_type="application/json", + ) + self.assertEqual(response.status_code, 200, response.content) + + def names(self): + """Return the sorted interface names of the module in each bay, on the active branch or on main.""" + return {position: sorted(name for _, name in self.interfaces_at(position)) for position in POSITIONS} + + def branch_names(self): + with self.in_branch(): + return self.names() + + +class _MergeAndRevertTests: + """The rule exists on main only, so the branch keeps NetBox's names, and a rename trigger on main would not.""" + + STRATEGY = None + + def setUp(self): + super().setUp() + self.add_rule() + self.install_and_move() + + def test_the_merge_gives_main_the_names_of_the_branch_and_only_the_replayed_changes(self): + job = self.enqueue("merge", **merging(self.STRATEGY)) + since = self.last_change_on_main() + + job = self.run_queued(job) + + self.assertEqual((job.status, job.error), (COMPLETED, "")) + self.assertEqual(self.branch_names(), BRANCH_NAMES) + self.assertEqual(self.names(), BRANCH_NAMES) + self.assert_only_replayed_changes_on_main(since) + + def test_the_revert_of_the_merge_gives_main_the_names_from_before_the_branch(self): + self.assertEqual(self.names(), NAMES_BEFORE) + self.assertEqual(self.act("merge", **merging(self.STRATEGY)).status, COMPLETED) + job = self.enqueue("revert", commit="on") + since = self.last_change_on_main() + + job = self.run_queued(job) + + self.assertEqual((job.status, job.error), (COMPLETED, "")) + self.assertEqual(self.names(), NAMES_BEFORE) + self.assert_only_replayed_changes_on_main(since) + + +class IterativeMergeTest(_MergeAndRevertTests, _InstallAndMoveCase): + PREFIX = "BrMergeIter" + STRATEGY = ITERATIVE + + +class SquashMergeTest(_MergeAndRevertTests, _InstallAndMoveCase): + PREFIX = "BrMergeSquash" + STRATEGY = SQUASH + + +class SyncTest(_InstallAndMoveCase): + """The rule exists in the branch only, so main keeps NetBox's names, and a rename trigger in the branch would not.""" + + PREFIX = "BrSync" + + def test_a_sync_writes_no_unlogged_rename_into_the_branch(self): + with self.in_branch(): + self.add_rule() + self.on_main() + self.install_and_move() + on_main = self.names() + + job = self.act("sync", commit="on") + + self.assertEqual((job.status, job.error), (COMPLETED, "")) + self.assertEqual(on_main, BRANCH_NAMES) + self.assertEqual(self.branch_names(), on_main) + + +class ReplayInTheActiveBranchTest(_ReplayCase): + """A merge and a revert started from a shell in which the branch is active. + + The branch changes the virtual-chassis position of a device, which is a rename trigger. While the + branch is active, netbox-branching itself refuses a replayed create and the revert of a module move. + """ + + PREFIX = "BrReplayActive" + + def build(self): + chassis = VirtualChassis.objects.create(name=f"{self.PREFIX} chassis") + device_type = make_device_type(make_manufacturer(self.PREFIX), self.PREFIX) + self.device = make_device(self.PREFIX, device_type, virtual_chassis=chassis, vc_position=1) + + def position_on_main(self): + return Device.objects.get(pk=self.device.pk).vc_position + + def test_a_merge_and_a_revert_started_in_the_branch_replay_a_rename_trigger_without_an_error(self): + response = self.client.patch( + reverse("dcim-api:device-detail", kwargs={"pk": self.device.pk}), + {"vc_position": 2}, + content_type="application/json", + ) + self.assertEqual(response.status_code, 200, response.content) + self.branch.refresh_from_db() + + with self.in_branch(): + self.branch.merge(user=self.user) + merged = self.position_on_main() + with self.in_branch(): + self.branch.revert(user=self.user) + + self.assertEqual((merged, self.position_on_main()), (2, 1)) + + +class RuleReplayInTheActiveBranchTest(_ReplayCase): + """A merge and a revert, started from a shell in which the branch is active, replay a rule update.""" + + PREFIX = "BrRuleReplayActive" + + def build(self): + module_type = make_module_type(make_manufacturer(self.PREFIX), self.PREFIX) + self.rule = InterfaceNameRule.objects.create(module_type=module_type, name_template="xe-{bay_position}") + + def template_on_main(self): + return InterfaceNameRule.objects.get(pk=self.rule.pk).name_template + + def test_a_merge_and_a_revert_started_in_the_branch_replay_a_rule_update_without_an_error(self): + response = self.client.patch( + reverse( + "plugins-api:netbox_interface_name_rules-api:interfacenamerule-detail", kwargs={"pk": self.rule.pk} + ), + {"name_template": "xe-0/{bay_position}", "description": "branch"}, + content_type="application/json", + ) + self.assertEqual(response.status_code, 200, response.content) + self.branch.refresh_from_db() + + with self.in_branch(): + self.branch.merge(user=self.user) + merged = self.template_on_main() + with self.in_branch(): + self.branch.revert(user=self.user) + + self.assertEqual((merged, self.template_on_main()), ("xe-0/{bay_position}", "xe-{bay_position}")) + + +class MergeExitTest(_InstallAndMoveCase): + """A merge that returns early or fails leaves the next save on main a rename trigger. + + The rule exists on main and in the branch, so an install on main gets its names only from a rename + trigger. + """ + + PREFIX = "BrMergeExit" + + def build(self): + super().build() + self.add_rule() + + def assert_an_install_on_main_gets_the_names_of_the_rule(self): + with transaction.atomic(): + Module.objects.create(device=self.device, module_bay=self.bay(0), module_type=self.module_type) + + self.assertEqual(self.names()[0], ["a0.br", "b0.br"]) + + def test_a_merge_without_a_change_leaves_the_next_save_a_rename_trigger(self): + job = self.act("merge", **merging(ITERATIVE)) + + self.assertEqual((job.status, job.error), (COMPLETED, "")) + self.assert_an_install_on_main_gets_the_names_of_the_rule() + + def test_a_dry_run_merge_leaves_the_next_save_a_rename_trigger(self): + self.install_and_move() + + job = self.act("merge", **merging(ITERATIVE, commit=False)) + + self.assertEqual((job.status, job.error), (COMPLETED, "")) + self.assertEqual(self.names(), NAMES_BEFORE) + self.assert_an_install_on_main_gets_the_names_of_the_rule() + + def test_a_failed_merge_leaves_the_next_save_a_rename_trigger(self): + self.install_and_move() + # The merge replays the creation of interface a0, and the device on main has an interface of that name now. + Interface.objects.create(device=self.device, name="a0", type=PLAIN_TYPE) + + job = self.act("merge", **merging(ITERATIVE)) + + self.assertEqual(job.status, ERRORED) + self.assertEqual(self.names(), NAMES_BEFORE) + Interface.objects.filter(device=self.device, name="a0").delete() + self.assert_an_install_on_main_gets_the_names_of_the_rule() + + def test_a_merge_of_the_branch_in_another_worker_does_not_skip_a_save_here(self): + """The replay mark belongs to the context that ran the operation, not to the branch status that workers share.""" + from netbox_branching.choices import BranchStatusChoices + from netbox_branching.models import Branch + + self.install_and_move() + self.assertEqual(self.act("merge", **merging(ITERATIVE, commit=False)).status, COMPLETED) + # netbox-branching sets this status first when another worker starts to merge the branch. + Branch.objects.filter(pk=self.branch.pk).update(status=BranchStatusChoices.MERGING) + + self.assert_an_install_on_main_gets_the_names_of_the_rule() + + +class _CascadeCase(KeptChannelCase, _ReplayCase): + """NetBox's cascade renames a kept channel again at the commit of a replay: the documented limit. + + The rule keeps channel 2 at ``1:2`` while it renames the parent. A replay renames the parent before + its commit, so NetBox's cascade gives channel 2 the name ``et-0/0/1:2`` at that commit. + """ + + def apply_on_the_page(self): + response = self.client.post( + self.apply_url(self.rule), {"action": "apply", "interface_ids": [str(self.parent("1").pk)]} + ) + self.assertEqual(response.status_code, 302) + + def names_on(self, alias): + return sorted(Interface.objects.using(alias).filter(module=self.modules["1"]).values_list("name", flat=True)) + + +class _CascadeMergeTests: + """A merge with one strategy, and the revert of that merge.""" + + STRATEGY = None + + def test_a_merge_renames_the_kept_channel_on_main_and_a_revert_restores_the_names_from_before(self): + self.apply_on_the_page() + self.assertEqual(self.names_on(self.alias), self.kept("1")) + + merged = self.act("merge", **merging(self.STRATEGY)) + names_after_the_merge = self.names_on("default") + reverted = self.act("revert", commit="on") + + self.assertEqual((merged.status, reverted.status), (COMPLETED, COMPLETED)) + self.assertEqual(names_after_the_merge, self.cascaded("1")) + self.assertEqual(self.names_on(self.alias), self.kept("1")) + self.assertEqual(self.names_on("default"), RAW_FAMILY) + + +class IterativeCascadeTest(_CascadeMergeTests, _CascadeCase): + PREFIX = "BrCascadeIter" + STRATEGY = ITERATIVE + + +class SquashCascadeTest(_CascadeMergeTests, _CascadeCase): + PREFIX = "BrCascadeSquash" + STRATEGY = SQUASH + + +class SyncCascadeTest(_CascadeCase): + PREFIX = "BrCascadeSync" + + def test_a_sync_renames_the_kept_channel_in_the_branch(self): + self.on_main() + self.apply_on_the_page() + self.assertEqual(self.names_on("default"), self.kept("1")) + + job = self.act("sync", commit="on") + + self.assertEqual((job.status, job.error), (COMPLETED, "")) + self.assertEqual(self.names_on(self.alias), self.cascaded("1")) + self.assertEqual(self.names_on("default"), self.kept("1")) diff --git a/netbox_interface_name_rules/tests/test_branch_transactions.py b/netbox_interface_name_rules/tests/test_branch_transactions.py new file mode 100644 index 00000000..5860403e --- /dev/null +++ b/netbox_interface_name_rules/tests/test_branch_transactions.py @@ -0,0 +1,351 @@ +# SPDX-License-Identifier: Apache-2.0 +# Copyright (C) 2025 Marcin Zieba +"""The plugin's write scope and blocks across ``default`` and the connection of a real netbox-branching branch.""" + +from contextlib import ExitStack + +from dcim.models import Interface, Site +from django.conf import settings +from django.contrib.auth import get_user_model +from django.db import DataError, IntegrityError, InternalError, OperationalError, connections, router, transaction +from django.test import override_settings +from extras.models import SavedFilter, Webhook +from netbox.context import events_queue + +from netbox_interface_name_rules.models import InterfaceNameRule +from netbox_interface_name_rules.tests.branch_cases import BRANCH_BEFORE, DEFAULT_BEFORE, SERVER_DEFAULT, BranchTestCase +from netbox_interface_name_rules.tests.helpers import ( + activate, + lock_timeout, + make_manufacturer, + make_module_type, + request_context, + set_lock_timeout, +) +from netbox_interface_name_rules.transactions import ( + LOCK_TIMEOUT, + SET_LOCK_TIMEOUT, + atomic_with_events, + on_commit, + write_scope, +) + +User = get_user_model() + + +def abort_the_transaction(alias): + """Make PostgreSQL refuse every later statement of the open transaction on *alias*.""" + with connections[alias].cursor() as cursor: + try: + cursor.execute("SELECT 1 / 0") + except DataError: + return + raise AssertionError("the statement did not fail") + + +def refuse_the_restore(execute, sql, params, many, context): + """Fail the statement that gives a connection its earlier ``lock_timeout`` back, before PostgreSQL runs it.""" + if sql == SET_LOCK_TIMEOUT and params != [LOCK_TIMEOUT]: + raise OperationalError("injected restore failure") + return execute(sql, params, many, context) + + +class _ScopeCase(BranchTestCase): + """One branch, and a distinct ``lock_timeout`` on each connection.""" + + def setUp(self): + user = User.objects.create_user(username=f"{type(self).__name__.lower()}-operator") + self.branch = self.provision_branch(type(self).__name__[:40], user) + self.alias = self.branch.connection_name + self.addCleanup(set_lock_timeout, "default", SERVER_DEFAULT) + set_lock_timeout("default", DEFAULT_BEFORE) + set_lock_timeout(self.alias, BRANCH_BEFORE) + + def alias_of(self, name): + return self.alias if name == "branch" else "default" + + def timeouts(self): + return lock_timeout("default"), lock_timeout(self.alias) + + +class WriteScopeInABranchTest(_ScopeCase): + def test_the_scope_pins_default_and_the_branch_alias(self): + with activate(self.branch), write_scope() as aliases: + self.assertEqual(aliases, ("default", self.alias)) + + def test_an_unexpected_write_alias_raises_before_any_query(self): + with activate(self.branch), self.assertNumQueries(0), self.assertNumQueries(0, using=self.alias): + with self.assertRaisesMessage(RuntimeError, f"'{self.alias}'"), write_scope(expected_alias="default"): + self.fail("the scope opened") + + def test_a_nested_scope_on_another_alias_raises(self): + with write_scope(), activate(self.branch), self.assertRaisesMessage(RuntimeError, f"'{self.alias}'"): + with write_scope(): + self.fail("the nested scope opened") + with activate(self.branch), write_scope(), activate(None), self.assertRaisesMessage(RuntimeError, "'default'"): + with write_scope(): + self.fail("the nested scope opened") + + +class LockTimeoutScopeTest(_ScopeCase): + def test_the_scope_sets_the_timeout_on_both_connections_and_restores_each_value(self): + with activate(self.branch): + with write_scope(): + self.assertEqual(self.timeouts(), (LOCK_TIMEOUT, LOCK_TIMEOUT)) + self.assertEqual(self.timeouts(), (DEFAULT_BEFORE, BRANCH_BEFORE)) + + def test_a_block_opens_the_scope_and_restores_after_its_commit_callbacks(self): + seen = [] + with activate(self.branch): + with atomic_with_events(): + on_commit(lambda: seen.append(self.timeouts())) + self.assertEqual(self.timeouts(), (DEFAULT_BEFORE, BRANCH_BEFORE)) + + self.assertEqual(seen, [(LOCK_TIMEOUT, LOCK_TIMEOUT)]) + + def test_a_nested_scope_does_not_restore(self): + with activate(self.branch), write_scope(): + with self.assertNumQueries(0), self.assertNumQueries(0, using=self.alias), write_scope(): + pass + self.assertEqual(self.timeouts(), (LOCK_TIMEOUT, LOCK_TIMEOUT)) + + def test_on_main_the_scope_changes_nothing(self): + with write_scope(): + self.assertEqual(lock_timeout("default"), DEFAULT_BEFORE) + + def test_a_restore_that_fails_discards_its_connection_and_the_other_is_still_restored(self): + default = connections["default"] + with activate(self.branch), default.execute_wrapper(refuse_the_restore): + with self.assertRaisesMessage(RuntimeError, "'default'") as raised, write_scope(): + session = default.connection + + self.assertEqual(str(raised.exception.__cause__), "injected restore failure") + self.assertTrue(session.closed) + self.assertIsNone(default.connection) + self.assertEqual(self.timeouts(), (SERVER_DEFAULT, BRANCH_BEFORE)) + + def test_inside_a_caller_transaction_the_set_and_the_restore_run_in_it(self): + for rolls_back in (False, True): + with self.subTest(rolls_back=rolls_back), activate(self.branch): + with transaction.atomic(using="default"), transaction.atomic(using=self.alias): + with write_scope(): + self.assertEqual(self.timeouts(), (LOCK_TIMEOUT, LOCK_TIMEOUT)) + self.assertEqual(self.timeouts(), (DEFAULT_BEFORE, BRANCH_BEFORE)) + if rolls_back: + transaction.set_rollback(True, using="default") + transaction.set_rollback(True, using=self.alias) + self.assertEqual(self.timeouts(), (DEFAULT_BEFORE, BRANCH_BEFORE)) + + def test_a_restore_that_fails_inside_a_caller_transaction_discards_its_connection(self): + branch_connection = connections[self.alias] + with activate(self.branch), self.assertRaisesMessage(RuntimeError, f"'{self.alias}'"): + with transaction.atomic(using=self.alias), branch_connection.execute_wrapper(refuse_the_restore): + with write_scope(): + session = branch_connection.connection + + self.assertTrue(session.closed) + self.assertIsNone(branch_connection.connection) + self.assertEqual(self.timeouts(), (DEFAULT_BEFORE, SERVER_DEFAULT)) + + def test_after_an_aborted_caller_transaction_no_session_keeps_the_timeout(self): + """PostgreSQL refuses the restore in the aborted transaction, so the scope discards that session.""" + with activate(self.branch), self.assertRaisesMessage(RuntimeError, "'default'"): + with transaction.atomic(using="default"), write_scope(): + abort_the_transaction("default") + + self.assertIsNone(connections["default"].connection) + self.assertEqual(self.timeouts(), (SERVER_DEFAULT, BRANCH_BEFORE)) + + def test_a_setup_that_fails_on_the_branch_restores_default(self): + with activate(self.branch), transaction.atomic(using=self.alias): + abort_the_transaction(self.alias) + with self.assertRaises(InternalError), write_scope(): + self.fail("the scope opened") + self.assertEqual(lock_timeout("default"), DEFAULT_BEFORE) + transaction.set_rollback(True, using=self.alias) + + +class CommitJoinTest(_ScopeCase): + """``on_commit`` runs its callback once, after the transactions open at the call commit on both connections.""" + + NESTINGS = ((), ("default",), ("branch",), ("default", "branch"), ("branch", "default")) + + def _in_transaction(self): + return connections["default"].in_atomic_block, connections[self.alias].in_atomic_block + + def _join(self, calls): + on_commit(lambda: calls.append(self._in_transaction())) + + def _register(self, calls): + with atomic_with_events(): + self._join(calls) + + def test_the_callback_runs_once_after_both_connections_commit_in_every_nesting(self): + for nesting in self.NESTINGS: + calls = [] + with self.subTest(nesting=nesting), activate(self.branch): + with ExitStack() as caller: + for name in nesting: + caller.enter_context(transaction.atomic(using=self.alias_of(name))) + self._register(calls) + if nesting: + self.assertEqual(calls, []) + self.assertEqual(calls, [(False, False)]) + + def test_a_connection_in_autocommit_acknowledges_at_once_and_the_callback_waits_for_the_other(self): + for held in ("default", "branch"): + calls = [] + with self.subTest(held=held), activate(self.branch): + with transaction.atomic(using=self.alias_of(held)), write_scope(): + self._join(calls) + self.assertEqual(calls, []) + self.assertEqual(calls, [(False, False)]) + + def test_a_rollback_of_either_transaction_drops_the_callback(self): + for nesting in (("default", "branch"), ("branch", "default")): + for rolled_back in nesting: + calls = [] + with self.subTest(nesting=nesting, rolled_back=rolled_back), activate(self.branch): + with ExitStack() as caller: + for name in nesting: + caller.enter_context(transaction.atomic(using=self.alias_of(name))) + self._register(calls) + transaction.set_rollback(True, using=self.alias_of(rolled_back)) + self.assertEqual(calls, []) + + def test_a_savepoint_rollback_around_the_call_on_either_connection_drops_the_callback(self): + for rolled_back in ("default", "branch"): + calls = [] + with self.subTest(rolled_back=rolled_back), activate(self.branch): + with transaction.atomic(using="default"), transaction.atomic(using=self.alias): + with transaction.atomic(using=self.alias_of(rolled_back)): + self._register(calls) + transaction.set_rollback(True, using=self.alias_of(rolled_back)) + self.assertEqual(calls, []) + + +class TwoAliasBlockTest(_ScopeCase): + """A block opens ``default`` outside the branch, and a nested block opens a savepoint on each.""" + + def _site(self, name): + return Site.objects.create(name=f"TwoAlias {name}", slug=f"twoalias-{name}") + + @staticmethod + def _webhook(name): + # netbox-branching exempts webhooks, so they stay on default in a branch. + return Webhook.objects.create(name=f"TwoAlias {name}", payload_url="http://localhost/") + + def test_the_branch_commits_before_default(self): + commits = [] + with activate(self.branch), atomic_with_events(): + transaction.on_commit(lambda: commits.append("default"), using="default") + transaction.on_commit(lambda: commits.append("branch"), using=self.alias) + + self.assertEqual(commits, ["branch", "default"]) + + def test_a_block_marked_for_rollback_writes_nothing_on_either_connection(self): + with activate(self.branch): + with atomic_with_events() as block: + self.assertEqual(block.aliases, ("default", self.alias)) + site, webhook = self._site("marked"), self._webhook("marked") + block.set_rollback() + self.assertFalse(Site.objects.filter(pk=site.pk).exists()) + self.assertFalse(Webhook.objects.filter(pk=webhook.pk).exists()) + + def test_a_nested_block_that_raises_rolls_back_only_its_own_writes_on_both_connections(self): + with activate(self.branch): + with atomic_with_events(): + kept = (self._site("kept"), self._webhook("kept")) + with self.assertRaises(ZeroDivisionError), atomic_with_events(): + dropped = (self._site("dropped"), self._webhook("dropped")) + _ = 1 / 0 + sites = set(Site.objects.filter(pk__in=(kept[0].pk, dropped[0].pk)).values_list("pk", flat=True)) + webhooks = set(Webhook.objects.filter(pk__in=(kept[1].pk, dropped[1].pk)).values_list("pk", flat=True)) + + self.assertEqual((sites, webhooks), ({kept[0].pk}, {kept[1].pk})) + + def test_a_default_commit_that_fails_after_the_branch_commit_keeps_the_events(self): + """The branch rows are committed, so the events of the block stay queued (the failure table).""" + user = User.objects.create_user(username="twoalias-events") + with activate(self.branch), request_context(user): + with self.assertRaises(IntegrityError) as raised, atomic_with_events(): + site = self._site("committed") + # No user has this ID; PostgreSQL checks the deferred foreign key at the COMMIT of default. + SavedFilter.objects.create( + name="TwoAlias dangling", slug="twoalias-dangling", user_id=2_000_000_000, parameters={} + ) + queued = dict(events_queue.get()) + self.assertTrue(Site.objects.filter(pk=site.pk).exists()) + + self.assertIn("foreign key", str(raised.exception)) + self.assertIn(f"dcim.site:{site.pk}", queued) + + +class RuleSaveInABranchTest(_ScopeCase): + """A rule save writes through the alias that the router gives for the rule, not through another alias.""" + + def test_a_save_through_default_in_a_branch_is_refused_before_any_query(self): + module_type = make_module_type(make_manufacturer("BrRuleSave"), "BrRuleSave") + rule = InterfaceNameRule.objects.create(module_type=module_type, name_template="xe-{bay_position}") + rule.name_template = "xe-0/{bay_position}" + + with activate(self.branch), self.assertNumQueries(0), self.assertNumQueries(0, using=self.alias): + with self.assertRaisesMessage(RuntimeError, f"'{self.alias}'"): + rule.save(using="default", update_fields=["name_template"]) + + self.assertEqual(InterfaceNameRule.objects.get(pk=rule.pk).name_template, "xe-{bay_position}") + + def test_a_full_save_through_default_in_a_branch_is_refused_before_any_query(self): + module_type = make_module_type(make_manufacturer("BrRuleFull"), "BrRuleFull") + rule = InterfaceNameRule.objects.create(module_type=module_type, name_template="xe-{bay_position}") + rule.name_template = "xe-0/{bay_position}" + + with activate(self.branch), self.assertNumQueries(0), self.assertNumQueries(0, using=self.alias): + with self.assertRaisesMessage(RuntimeError, f"'{self.alias}'"): + rule.save(using="default") + + self.assertEqual(InterfaceNameRule.objects.get(pk=rule.pk).name_template, "xe-{bay_position}") + + def test_a_description_save_through_default_in_a_branch_is_refused_before_any_query(self): + module_type = make_module_type(make_manufacturer("BrRuleDesc"), "BrRuleDesc") + rule = InterfaceNameRule.objects.create(module_type=module_type, name_template="xe-{bay_position}") + rule.description = "after" + + with activate(self.branch), self.assertNumQueries(0), self.assertNumQueries(0, using=self.alias): + with self.assertRaisesMessage(RuntimeError, f"'{self.alias}'"): + rule.save(using="default", update_fields=["description"]) + + self.assertEqual(InterfaceNameRule.objects.get(pk=rule.pk).description, "") + + +class ExemptRuleModelInABranchTest(_ScopeCase): + """An operator can exempt the plugin's models from netbox-branching; a rule then lives on ``default`` alone.""" + + def setUp(self): + module_type = make_module_type(make_manufacturer("BrExempt"), "BrExempt") + # Made before the branch, so the branch holds a copy that the save must leave alone. + self.rule = InterfaceNameRule.objects.create(module_type=module_type, name_template="xe-{bay_position}") + super().setUp() + branching = {**settings.PLUGINS_CONFIG["netbox_branching"], "exempt_models": ["netbox_interface_name_rules.*"]} + exempt = override_settings(PLUGINS_CONFIG={**settings.PLUGINS_CONFIG, "netbox_branching": branching}) + exempt.enable() + self.addCleanup(exempt.disable) + + def test_the_router_gives_default_for_the_rule_and_the_branch_for_an_interface(self): + with activate(self.branch): + routed = router.db_for_write(InterfaceNameRule), router.db_for_write(Interface) + + self.assertEqual(routed, ("default", self.alias)) + + def test_a_partial_rule_save_in_a_branch_writes_on_default(self): + self.rule.name_template = "xe-0/{bay_position}" + + with activate(self.branch): + self.rule.save(update_fields=["name_template"]) + + self.assertEqual( + InterfaceNameRule.objects.using("default").get(pk=self.rule.pk).name_template, "xe-0/{bay_position}" + ) + self.assertEqual( + InterfaceNameRule.objects.using(self.alias).get(pk=self.rule.pk).name_template, "xe-{bay_position}" + ) diff --git a/netbox_interface_name_rules/tests/test_branch_triggers.py b/netbox_interface_name_rules/tests/test_branch_triggers.py new file mode 100644 index 00000000..f5d39a49 --- /dev/null +++ b/netbox_interface_name_rules/tests/test_branch_triggers.py @@ -0,0 +1,321 @@ +# SPDX-License-Identifier: Apache-2.0 +# Copyright (C) 2025 Marcin Zieba +"""Rename triggers in a real netbox-branching branch: the plan runs after the commit of the branch connection. + +NetBox saves a new module before it creates the module's interfaces, in the transaction that it opens +on the branch connection. Each test builds its rows on main first, because a branch copies main when +it provisions. Requests carry netbox-branching's cookie or header, so its request processor, NetBox's +view transaction, the trigger and the plan run for real. +""" + +import uuid +from contextlib import ExitStack + +from core.choices import ObjectChangeActionChoices +from core.models import ObjectChange +from dcim.models import Device, Interface, InterfaceTemplate, Module +from django.contrib.contenttypes.models import ContentType +from django.db import connections, router, transaction +from django.db.models.signals import post_save +from django.test import RequestFactory +from django.urls import reverse +from extras.choices import JournalEntryKindChoices +from extras.jobs import ScriptJob +from extras.models import JournalEntry +from extras.scripts import Script +from netbox.registry import registry + +from netbox_interface_name_rules.models import InterfaceNameRule +from netbox_interface_name_rules.rename_triggers import PlanRunner +from netbox_interface_name_rules.tests.branch_cases import ( + BRANCH_BEFORE, + DEFAULT_BEFORE, + SERVER_DEFAULT, + BranchWriteCase, + ChannelCase, +) +from netbox_interface_name_rules.tests.helpers import ( + PLAIN_TYPE, + branch_cookie, + install_form, + interface_signal, + lock_timeout, + make_device, + make_device_type, + make_job, + make_manufacturer, + make_module_bay_templates, + make_module_type, + names_of, + set_lock_timeout, +) +from netbox_interface_name_rules.transactions import LOCK_TIMEOUT + +# The caller's transactions at the save, outermost first, when the branch connection holds one. +NESTINGS = (("branch",), ("default", "branch"), ("branch", "default")) + + +class _InstallCase(BranchWriteCase): + """A device with empty module bays, and a module type whose rule appends ``.br`` to NetBox's names. + + NetBox names the two interfaces of a module in the bay at position ``n`` ``an`` and ``bn``. + """ + + BAYS = 3 + + def build(self): + manufacturer = make_manufacturer(self.PREFIX) + device_type = make_device_type(manufacturer, self.PREFIX) + make_module_bay_templates(device_type, tuple(f"Bay {position}" for position in range(self.BAYS))) + self.device = make_device(self.PREFIX, device_type) + self.module_type = make_module_type(manufacturer, self.PREFIX) + for template in ("a{module}", "b{module}"): + InterfaceTemplate.objects.create(module_type=self.module_type, name=template, type=PLAIN_TYPE) + InterfaceNameRule.objects.create(module_type=self.module_type, name_template="{base}.br") + + def install(self, position): + """Install a module in the bay at *position* through the ORM, on the alias that the router gives.""" + return Module.objects.create(device=self.device, module_bay=self.bay(position), module_type=self.module_type) + + def renames_in_branch(self, position): + """Return ``(name before, name after)`` of each interface update in the branch at *position*, sorted.""" + changes = ObjectChange.objects.using(self.alias).filter( + changed_object_type=ContentType.objects.get_for_model(Interface), + changed_object_id__in=[pk for pk, _ in self.interfaces_in_branch(position)], + action=ObjectChangeActionChoices.ACTION_UPDATE, + ) + return sorted((change.prechange_data["name"], change.postchange_data["name"]) for change in changes) + + def assert_nothing_on_main(self): + self.assertFalse(Interface.objects.filter(device=self.device).exists()) + + +class InstallInABranchTest(_InstallCase): + PREFIX = "BrInstall" + + def test_a_module_installed_through_the_ui_gets_its_names_in_the_branch(self): + response = self.client.post(reverse("dcim:module_add"), install_form(self.bay(0), self.module_type)) + + self.assertEqual(response.status_code, 302) + self.assertEqual(self.names_in_branch(0), ["a0.br", "b0.br"]) + self.assertEqual(self.renames_in_branch(0), [("a0", "a0.br"), ("b0", "b0.br")]) + self.assert_nothing_on_main() + + def test_a_module_installed_through_the_rest_api_gets_its_names_in_the_branch(self): + del self.client.cookies[branch_cookie()] + data = {"device": self.device.pk, "module_bay": self.bay(0).pk, "module_type": self.module_type.pk} + + response = self.client.post( + reverse("dcim-api:module-list"), + data, + content_type="application/json", + headers={"X-NetBox-Branch": self.branch.schema_id}, + ) + + self.assertEqual(response.status_code, 201, response.content) + self.assertEqual(self.names_in_branch(0), ["a0.br", "b0.br"]) + self.assertEqual(self.renames_in_branch(0), [("a0", "a0.br"), ("b0", "b0.br")]) + self.assert_nothing_on_main() + + def test_modules_imported_in_bulk_get_their_names_in_the_branch(self): + rows = [f"{self.device.name},Bay {position},{self.module_type.model},active" for position in (0, 1)] + + response = self.client.post( + reverse("dcim:module_bulk_import"), + { + "data": "\n".join(["device,module_bay,module_type,status", *rows]), + "format": "csv", + "csv_delimiter": "auto", + }, + ) + + self.assertEqual(response.status_code, 302) + self.assertEqual( + [self.names_in_branch(position) for position in (0, 1)], [["a0.br", "b0.br"], ["a1.br", "b1.br"]] + ) + self.assertEqual( + [self.renames_in_branch(position) for position in (0, 1)], + [[("a0", "a0.br"), ("b0", "b0.br")], [("a1", "a1.br"), ("b1", "b1.br")]], + ) + self.assert_nothing_on_main() + + def test_a_name_collision_on_install_skips_that_interface_and_keeps_the_rest_of_the_save(self): + with self.in_branch(): + Interface.objects.create(device=self.device, name="b0.br", type=PLAIN_TYPE) + + response = self.client.post(reverse("dcim:module_add"), install_form(self.bay(0), self.module_type)) + + self.assertEqual(response.status_code, 302) + self.assertEqual(self.names_in_branch(0), ["a0.br", "b0"]) + self.assertEqual(self.renames_in_branch(0), [("a0", "a0.br")]) + with self.in_branch(): + module = Module.objects.get(module_bay=self.bay(0)) + (entry,) = JournalEntry.objects.filter( + assigned_object_type=ContentType.objects.get_for_model(Module), assigned_object_id=module.pk + ) + self.assertEqual(entry.kind, JournalEntryKindChoices.KIND_WARNING) + self.assertIn("`b0` to `b0.br`", entry.comments) + self.assert_nothing_on_main() + + def test_a_script_install_in_the_transaction_of_the_configuration_guide_gets_its_names_in_the_branch(self): + with self.in_branch(), transaction.atomic(using=router.db_for_write(Interface)): + self.install(0) + + self.assertEqual(self.names_in_branch(0), ["a0.br", "b0.br"]) + self.assert_nothing_on_main() + + def test_a_script_install_in_a_transaction_on_default_alone_gets_no_names_in_the_branch(self): + with self.in_branch(), transaction.atomic(): + self.install(0) + + self.assertEqual(self.names_in_branch(0), ["a0", "b0"]) + + +class MoveInABranchTest(_InstallCase): + """A module installed on main, so the branch copies it with the names of the rule.""" + + PREFIX = "BrMove" + + def build(self): + super().build() + with transaction.atomic(): + self.module = self.install(0) + + def test_a_module_moved_through_the_rest_api_gets_the_names_of_its_new_bay_in_the_branch(self): + del self.client.cookies[branch_cookie()] + + response = self.client.patch( + reverse("dcim-api:module-detail", kwargs={"pk": self.module.pk}), + {"module_bay": self.bay(1).pk}, + content_type="application/json", + headers={"X-NetBox-Branch": self.branch.schema_id}, + ) + + self.assertEqual(response.status_code, 200, response.content) + self.assertEqual(self.names_in_branch(1), ["a1.br", "b1.br"]) + self.assertEqual(self.renames_in_branch(1), [("a0.br", "a1.br"), ("b0.br", "b1.br")]) + self.assertEqual(sorted(names_of(self.module)), ["a0.br", "b0.br"]) + self.assertEqual(Module.objects.get(pk=self.module.pk).module_bay, self.bay(0)) + + +class TriggerTransactionsInABranchTest(_InstallCase): + PREFIX = "BrTrigger" + + def test_trigger_savepoint_rollback_filters_only_rolled_back_triggers(self): + """A rollback of a branch savepoint drops its trigger; a rollback on ``default`` drops no branch save.""" + with self.in_branch(), transaction.atomic(using="default"), transaction.atomic(using=self.alias): + self.install(0) + with transaction.atomic(using=self.alias): + self.install(1) + transaction.set_rollback(True, using=self.alias) + with transaction.atomic(using="default"): + self.install(2) + transaction.set_rollback(True, using="default") + runners = [entry for _, entry, _ in connections[self.alias].run_on_commit if isinstance(entry, PlanRunner)] + + self.assertEqual( + [self.names_in_branch(position) for position in (0, 1, 2)], [["a0.br", "b0.br"], [], ["a2.br", "b2.br"]] + ) + self.assertEqual({id(runner.plan) for runner in runners}, {id(runners[-1].plan)}) + self.assertEqual([trigger.kept for trigger in runners[-1].plan.triggers], [True, False, True]) + + def test_a_plan_that_commits_after_its_branch_was_left_raises(self): + with self.assertRaisesMessage(RuntimeError, f"'default', but the operation expects '{self.alias}'"): + with transaction.atomic(using=self.alias), self.in_branch(): + self.install(0) + + self.assertEqual(self.names_in_branch(0), ["a0", "b0"]) + + def test_a_save_through_default_in_a_branch_raises_before_the_row_is_written(self): + with self.in_branch(): + device = Device.objects.get(pk=self.device.pk) + device.name = "brtrigger-renamed" + + with self.in_branch(), self.assertRaisesMessage(RuntimeError, f"but the write alias is '{self.alias}'"): + device.save(using="default") + + self.assertEqual(Device.objects.get(pk=self.device.pk).name, self.device.name) + + +class _RuledChannelCase(ChannelCase): + """The channelized case with its rule on main, so an install in the branch is a rename trigger.""" + + def build(self): + super().build() + self.add_rule() + + +class ChannelizedInstallInABranchTest(_RuledChannelCase): + PREFIX = "BrChanInstall" + POSITIONS = ("1", "2", "3", "4") + + def test_trigger_runner_reconciliation_during_branch_callback_drain(self): + """The plan runs in the commit callbacks of the branch; the reconciliation still follows NetBox's cascade.""" + response = self.client.post(reverse("dcim:module_add"), install_form(self.bay("1"), self.module_type)) + + self.assertEqual(response.status_code, 302) + self.assertEqual(self.names_in_branch("1"), self.kept("1")) + for position, nesting in zip(("2", "3", "4"), NESTINGS, strict=True): + with self.subTest(nesting=nesting), self.in_branch(): + with ExitStack() as caller: + for name in nesting: + caller.enter_context(transaction.atomic(using=self.alias if name == "branch" else "default")) + Module.objects.create( + device=self.device, module_bay=self.bay(position), module_type=self.module_type + ) + self.assertEqual(self.names_in_branch(position), self.kept(position)) + + +class ScriptTriggerInABranchTest(_RuledChannelCase): + """A NetBox script installs a module in a branch; each connection starts from its own ``lock_timeout``.""" + + PREFIX = "BrScript" + + def setUp(self): + super().setUp() + self.addCleanup(set_lock_timeout, "default", SERVER_DEFAULT) + set_lock_timeout("default", DEFAULT_BEFORE) + set_lock_timeout(self.alias, BRANCH_BEFORE) + + def timeouts(self): + return lock_timeout("default"), lock_timeout(self.alias) + + def run_script(self, script): + """Run *script* as NetBox's script job does, in the request processors of a request in the branch.""" + request = RequestFactory().get("/") + request.user = self.user + request.id = uuid.uuid4() + request.COOKIES[branch_cookie()] = self.branch.schema_id + with ExitStack() as processors: + for processor in registry["request_processors"]: + processors.enter_context(processor(request)) + ScriptJob(make_job(self.PREFIX, self.user)).run_script(script, request, {}, commit=True) + + def test_script_trigger_restores_timeouts_before_caller_default_callbacks(self): + seen = {} + case = self + bay = self.bay("1") + + class InstallModule(Script): + def run(self, data, commit): + transaction.on_commit( + lambda: seen.setdefault("script callback", (case.timeouts(), case.names_in_branch("1"))), + using="default", + ) + Module.objects.create(device_id=bay.device_id, module_bay_id=bay.pk, module_type=case.module_type) + + def at_the_cascade(sender, instance, **kwargs): + if instance.name == "et-0/0/1:2": + seen.setdefault("cascade", self.timeouts()) + + with interface_signal(post_save, at_the_cascade): + self.run_script(InstallModule()) + + self.assertEqual( + seen, + { + "cascade": (LOCK_TIMEOUT, LOCK_TIMEOUT), + "script callback": ((DEFAULT_BEFORE, BRANCH_BEFORE), self.cascaded("1")), + }, + ) + self.assertEqual(self.names_in_branch("1"), self.kept("1")) diff --git a/netbox_interface_name_rules/tests/test_branch_writes.py b/netbox_interface_name_rules/tests/test_branch_writes.py new file mode 100644 index 00000000..7ec13522 --- /dev/null +++ b/netbox_interface_name_rules/tests/test_branch_writes.py @@ -0,0 +1,361 @@ +# SPDX-License-Identifier: Apache-2.0 +# Copyright (C) 2025 Marcin Zieba +"""Plugin writes started in a real netbox-branching branch land in the branch, on both connections at once. + +Each test builds its rows on main first, because a branch copies main when it provisions. Requests +carry netbox-branching's cookie, so its request processor, NetBox's change log and the ChangeDiff rows +that netbox-branching writes on ``default`` run for real. +""" + +from concurrent.futures import ThreadPoolExecutor +from contextlib import ExitStack +from functools import partial +from unittest import skipUnless + +from core.models import ObjectChange +from dcim.models import Interface, VirtualChassis +from django.contrib.auth import get_user_model +from django.contrib.messages import get_messages +from django.db import connections, transaction +from django.db.models.signals import post_save, pre_save +from django.urls import reverse + +from netbox_interface_name_rules.engine import ( + apply_device_interface_rules, + find_matching_rule, + supports_channelization, +) +from netbox_interface_name_rules.models import InterfaceNameRule +from netbox_interface_name_rules.rule_selection import pinned_rule_cache +from netbox_interface_name_rules.tests.branch_cases import ( + FLAT_NAMES, + BranchWriteCase, + ConversionCase, + KeptChannelCase, + PlainModuleCase, +) +from netbox_interface_name_rules.tests.helpers import ( + CHANNEL_TYPE, + FLAT, + PARENT_TYPE, + PLAIN_TYPE, + REQUIRES_CHANNELIZATION, + activate, + interface_signal, + make_device, + make_device_type, + make_manufacturer, + make_placement, + names_of, + row_lock_in_another_session, + set_lock_timeout, +) +from netbox_interface_name_rules.tests.out_of_band import rename_out_of_band + +User = get_user_model() + + +def messages_of(response): + """Return ``(level tag, text)`` for each message that *response*'s request produced.""" + return [(message.level_tag, str(message)) for message in get_messages(response.wsgi_request)] + + +def _create_in_its_own_session(branch, fields): + try: + # A lock wait on the test's own transaction would otherwise hang the test. + set_lock_timeout(branch.connection_name, "5s") + with activate(branch): + return Interface.objects.create(**fields).pk + finally: + connections[branch.connection_name].close() + + +def create_in_another_session(branch, **fields): + """Create an interface in *branch* through a session of another thread, which commits at once.""" + with ThreadPoolExecutor(max_workers=1) as executor: + return executor.submit(_create_in_its_own_session, branch, fields).result() + + +class ForegroundApplyInABranchTest(PlainModuleCase): + PREFIX = "BrApply" + + def build(self): + super().build() + self.rule = InterfaceNameRule.objects.create( + module_type=self.module_type, name_template="et-0/0/{bay_position}" + ) + + def test_the_apply_renames_in_the_branch_only(self): + response = self.client.post( + self.apply_url(self.rule), {"action": "apply", "interface_ids": [str(self.interface.pk)]} + ) + + self.assertEqual(messages_of(response), [("success", "Applied rule: 1 interface(s) renamed.")]) + with self.in_branch(): + self.assertEqual(names_of(self.module), ["et-0/0/0"]) + self.assertEqual(names_of(self.module), ["0"]) + self.assertEqual(self.branch_updates_of(self.interface), [("0", "et-0/0/0")]) + self.assertEqual(list(self.change_diffs().values_list("object_id", flat=True)), [self.interface.pk]) + + def test_the_toggle_writes_the_flag_in_the_branch_only(self): + url = reverse("plugins:netbox_interface_name_rules:interfacenamerule_toggle", args=[self.rule.pk]) + + response = self.client.post(url, HTTP_X_REQUESTED_WITH="XMLHttpRequest") + + self.assertEqual(response.json(), {"enabled": False, "pk": self.rule.pk}) + with self.in_branch(): + self.assertFalse(InterfaceNameRule.objects.get(pk=self.rule.pk).enabled) + self.assertTrue(InterfaceNameRule.objects.get(pk=self.rule.pk).enabled) + + +class TwoAliasAtomicityTest(PlainModuleCase): + PREFIX = "BrAtomic" + + def build(self): + super().build() + # The eleventh name is one character longer than NetBox allows, after ten rows were written. + self.rule = InterfaceNameRule.objects.create( + module_type=self.module_type, + name_template=f"{'e' * 60}-{{bay_position}}:{{channel}}", + breakout_mode=FLAT, + channel_count=11, + channel_start=0, + ) + + def test_a_flat_family_blocked_after_partial_writes_leaves_no_row_and_no_change_diff(self): + response = self.client.post( + self.apply_url(self.rule), {"action": "apply", "interface_ids": [str(self.interface.pk)]} + ) + + self.assertEqual( + messages_of(response), + [ + ("success", "Applied rule: 0 interface(s) renamed."), + ("warning", "1 interface(s) skipped. The plugin log names each one."), + ], + ) + with self.in_branch(): + self.assertEqual(names_of(self.module), ["0"]) + self.assertEqual(ObjectChange.objects.using(self.alias).filter(changed_object_id=self.interface.pk).count(), 0) + self.assertFalse(self.change_diffs().exists()) + + +class RuleCacheInABranchTest(PlainModuleCase): + PREFIX = "BrCache" + + def build(self): + super().build() + self.rule = InterfaceNameRule.objects.create( + module_type=self.module_type, name_template="et-0/0/{bay_position}" + ) + + def _select(self): + return find_matching_rule(self.module_type, None, self.device_type) + + def test_each_alias_gets_its_own_rule_instance(self): + on_main = self._select() + with self.in_branch(): + in_branch = self._select() + + self.assertEqual((on_main.pk, on_main._state.db), (self.rule.pk, "default")) + self.assertEqual((in_branch.pk, in_branch._state.db), (self.rule.pk, self.alias)) + + def test_a_pinned_cache_does_not_serve_main_in_the_branch(self): + with pinned_rule_cache(): + self._select() + with self.in_branch(): + self.assertEqual(self._select()._state.db, self.alias) + + def test_a_rule_changed_in_the_branch_only_changes_the_selection_in_the_branch_only(self): + self._select() + with self.in_branch(): + self._select() + rule = InterfaceNameRule.objects.get(pk=self.rule.pk) + rule.snapshot() + rule.name_template = "xe-0/0/{bay_position}" + rule.save() + in_branch = self._select() + on_main = self._select() + + self.assertEqual((in_branch.name_template, in_branch._state.db), ("xe-0/0/{bay_position}", self.alias)) + self.assertEqual((on_main.name_template, on_main._state.db), ("et-0/0/{bay_position}", "default")) + + +class ForegroundConvertInABranchTest(ConversionCase): + PREFIX = "BrConvert" + + def test_the_conversion_rewrites_the_family_in_the_branch_only(self): + response = self.client.post( + self.apply_url(self.rule), {"action": "convert", "convert_ids": [str(self.base.pk)]} + ) + + self.assertEqual( + messages_of(response), [("success", "Converted 1 interface family(ies) to the channelized topology.")] + ) + with self.in_branch(): + self.assertEqual(names_of(self.module), ["et-0/0/3", *FLAT_NAMES]) + self.assertEqual(names_of(self.module), list(FLAT_NAMES)) + + def test_the_conversion_preview_writes_nothing_on_either_connection(self): + response = self.client.get(self.apply_url(self.rule)) + + self.assertEqual([verdict.current_name for verdict in response.context["conversions"]], ["xe-0/0/3:0"]) + with self.in_branch(): + self.assertEqual(names_of(self.module), list(FLAT_NAMES)) + self.assertFalse(Interface.objects.filter(module=self.module, channels__isnull=False).exists()) + self.assertFalse(ObjectChange.objects.using(self.alias).exists()) + self.assertFalse(self.change_diffs().exists()) + + +class CommitOrderInEveryNestingTest(KeptChannelCase): + PREFIX = "BrNesting" + POSITIONS = ("1", "2", "3", "4", "5") + # The caller's transactions at the call, outermost first. + NESTINGS = ( + ("1", ()), + ("2", ("default",)), + ("3", ("branch",)), + ("4", ("default", "branch")), + ("5", ("branch", "default")), + ) + + def _enter(self, stack, nesting): + for name in nesting: + stack.enter_context(transaction.atomic(using=self.alias if name == "branch" else "default")) + + def test_on_commit_orders_cascade_and_reconciliation_for_each_transaction_nesting(self): + for position, nesting in self.NESTINGS: + with self.subTest(nesting=nesting), self.in_branch(): + with ExitStack() as caller: + self._enter(caller, nesting) + self.apply(position) + self.assertEqual(self.names_in_branch(position), self.kept(position)) + + def test_branch_outer_default_rollback_discards_reconciliation(self): + with self.in_branch(), transaction.atomic(using=self.alias), transaction.atomic(using="default"): + self.apply("1") + transaction.set_rollback(True, using="default") + + self.assertEqual(self.names_in_branch("1"), self.cascaded("1")) + + def test_default_savepoint_rollback_discards_reconciliation_before_branch_commit(self): + with self.in_branch(), transaction.atomic(using="default"), transaction.atomic(using=self.alias): + with transaction.atomic(using="default"): + self.apply("1") + transaction.set_rollback(True, using="default") + + self.assertEqual(self.names_in_branch("1"), self.cascaded("1")) + + +class CollisionInABranchTest(KeptChannelCase): + PREFIX = "BrCollide" + + def test_a_collision_via_apply_skips_that_member_and_keeps_the_rest(self): + """Another session takes the target of channel 3 after the plugin checked it, before the member saves.""" + channel = Interface.objects.get(module=self.modules["1"], channel_id=3) + occupied = [] + + def occupy(sender, instance, **kwargs): + if instance.pk == channel.pk and instance.name == "xe-0/0/1:2" and not occupied: + occupied.append( + create_in_another_session(self.branch, device_id=self.device.pk, name="xe-0/0/1:2", type=PLAIN_TYPE) + ) + + with interface_signal(pre_save, occupy): + response = self.client.post( + self.apply_url(self.rule), {"action": "apply", "interface_ids": [str(self.parent("1").pk)]} + ) + + self.assertEqual( + messages_of(response), + [ + ("success", "Applied rule: 3 interface(s) renamed."), + ("warning", "2 interface(s) skipped. The plugin log names each one."), + ], + ) + self.assertEqual(self.names_in_branch("1"), ["1:2", "1:3", "et-0/0/1", "xe-0/0/1:0", "xe-0/0/1:3"]) + with self.in_branch(): + self.assertEqual(Interface.objects.get(pk=occupied[0]).name, "xe-0/0/1:2") + + +class ReconciliationLockTimeoutTest(KeptChannelCase): + PREFIX = "BrReconcile" + + def test_a_reconciliation_lock_timeout_keeps_the_change_diff_rows_and_names_each_kept_channel(self): + """A second session locks the kept channel after NetBox's cascade renamed it, before the reconciliation.""" + channel = Interface.objects.get(module=self.modules["1"], channel_id=2) + + with row_lock_in_another_session(self.alias) as lock: + + def lock_after_the_cascade(sender, instance, **kwargs): + if instance.pk == channel.pk and instance.name == "et-0/0/1:2": + transaction.on_commit(partial(lock, channel.pk), using=instance._state.db) + + with interface_signal(post_save, lock_after_the_cascade): + response = self.client.post( + self.apply_url(self.rule), {"action": "apply", "interface_ids": [str(self.parent("1").pk)]} + ) + + [(level, text)] = messages_of(response) + self.assertEqual(level, "danger") + self.assertTrue(text.startswith(f"Failed to apply rule {self.rule}: "), text) + self.assertIn("Rename each channel back: `et-0/0/1:2` to `1:2`", text) + self.assertEqual(self.names_in_branch("1"), self.cascaded("1")) + with self.in_branch(): + renamed = set( + Interface.objects.filter(module=self.modules["1"]).exclude(pk=channel.pk).values_list("pk", flat=True) + ) + self.assertLessEqual(renamed, set(self.change_diffs().values_list("object_id", flat=True))) + + +@skipUnless(supports_channelization(), REQUIRES_CHANNELIZATION) +class DocumentedEngineFunctionInACallerTransactionTest(BranchWriteCase): + """``apply_device_interface_rules`` inside a caller's transaction keeps a blocked channel at its name. + + The caller frees the target of the blocked channel before it commits, so NetBox's cascade at the + commit gives the channel that target. The reconciliation runs after the cascade and gives the + channel back its kept name. + """ + + PREFIX = "BrEngine" + # The caller's transactions at the call, outermost first, one device each. + NESTINGS = (("branch",), ("default", "branch"), ("branch", "default")) + + def build(self): + device_type = make_device_type(make_manufacturer(self.PREFIX), self.PREFIX) + placement = make_placement(self.PREFIX) + chassis = VirtualChassis.objects.create(name=f"{self.PREFIX} chassis") + self.devices = [] + for position in range(1, len(self.NESTINGS) + 1): + device = make_device( + self.PREFIX, + device_type, + placement, + name=f"brengine-{position}", + virtual_chassis=chassis, + vc_position=position, + ) + parent = Interface.objects.create(device=device, name="et0", type=PARENT_TYPE, channels=4) + for channel_id in range(1, 5): + Interface.objects.create( + device=device, name=f"et0:{channel_id}", type=CHANNEL_TYPE, parent=parent, channel_id=channel_id + ) + # Channel 2 cannot take its target, so the rule keeps its name. + Interface.objects.create(device=device, name=f"eth{position}:2", type=PLAIN_TYPE) + self.devices.append(device) + InterfaceNameRule.objects.create( + applies_to_device_interfaces=True, module_type_pattern=r"et\d+", name_template="eth{vc_position}" + ) + + def test_a_blocked_channel_keeps_its_name_after_the_caller_frees_its_target_and_commits(self): + for position, (device, nesting) in enumerate(zip(self.devices, self.NESTINGS, strict=True), start=1): + with self.subTest(nesting=nesting), self.in_branch(): + with ExitStack() as caller: + for name in nesting: + caller.enter_context(transaction.atomic(using=self.alias if name == "branch" else "default")) + self.assertEqual(apply_device_interface_rules(device), 4) + rename_out_of_band(Interface.objects.get(device=device, name=f"eth{position}:2"), f"free{position}") + names = sorted(Interface.objects.filter(device=device).values_list("name", flat=True)) + self.assertEqual( + names, ["et0:2", f"eth{position}", *(f"eth{position}:{c}" for c in (1, 3, 4)), f"free{position}"] + ) diff --git a/netbox_interface_name_rules/tests/test_branching.py b/netbox_interface_name_rules/tests/test_branching.py index 354d0434..33682e0b 100644 --- a/netbox_interface_name_rules/tests/test_branching.py +++ b/netbox_interface_name_rules/tests/test_branching.py @@ -6,23 +6,27 @@ ``EXPECT_NETBOX_BRANCHING=1``, and there a missing netbox-branching fails the guard test instead. """ +import ast +import collections +import hashlib +import inspect import os import subprocess import sys +import tempfile +from pathlib import Path from unittest import skipUnless from dcim.models import Interface -from django.apps import apps from django.contrib.auth import get_user_model from django.core.exceptions import ImproperlyConfigured from django.db import connection, connections, router -from django.test import SimpleTestCase, TransactionTestCase +from django.test import SimpleTestCase -from netbox_interface_name_rules.branching import check_version -from netbox_interface_name_rules.tests.helpers import make_device, make_device_type, make_manufacturer - -BRANCHING_INSTALLED = apps.is_installed("netbox_branching") -BRANCHING_SKIP_REASON = "netbox-branching is not installed" +from netbox_interface_name_rules import branching +from netbox_interface_name_rules.branching import REPLAY_MARK, REPLAYING_METHODS, check_version +from netbox_interface_name_rules.tests.branch_cases import BRANCHING_INSTALLED, BRANCHING_SKIP_REASON, BranchTestCase +from netbox_interface_name_rules.tests.helpers import activate, make_device, make_device_type, make_manufacturer def schema_exists(schema_name): @@ -72,27 +76,301 @@ def test_an_unsupported_version_stops_startup(self): self.assertIn("is 1.3.0.", completed.stderr) -def remove_branch(branch): - """Close the connection of *branch*, then drop its schema.""" - connections[branch.connection_name].close() - branch.deprovision() +@skipUnless(BRANCHING_INSTALLED, BRANCHING_SKIP_REASON) +class ReplayWrapperContractTest(SimpleTestCase): + """The plugin wraps each method of netbox-branching's Branch that replays changes, once, as it was reviewed.""" + + def replaying_methods(self): + from netbox_branching.models import Branch + + return {name: getattr(Branch, name) for name in REPLAYING_METHODS} + + def test_each_replaying_method_is_wrapped_once_and_keeps_its_signature_and_attributes(self): + for name, method in self.replaying_methods().items(): + with self.subTest(method=name): + original = method.__wrapped__ + + self.assertFalse(getattr(original, REPLAY_MARK, False)) + self.assertEqual(str(inspect.signature(original)), "(self, user, commit=True)") + self.assertEqual(inspect.signature(method), inspect.signature(original)) + self.assertEqual((method.__name__, method.alters_data), (name, True)) + + def test_a_second_start_wraps_nothing_again(self): + wrapped = self.replaying_methods() + + branching.ready() + + self.assertEqual(self.replaying_methods(), wrapped) + + +# Each (name, receiver, module, scope) reference to a name that saves a replayed object in 1.2.1: count, what reaches it. +REPLAY_SAVES = { + ("apply", "change", "merge_strategies/iterative.py", "IterativeMergeStrategy.merge"): (1, "Branch.merge"), + ("undo", "change", "merge_strategies/iterative.py", "IterativeMergeStrategy.revert"): (1, "Branch.revert"), + ("apply", "dummy_change", "merge_strategies/squash.py", "SquashMergeStrategy.merge"): (1, "Branch.merge"), + ("undo", "dummy_change", "merge_strategies/squash.py", "SquashMergeStrategy.revert"): (1, "Branch.revert"), + ("apply", "change", "models/branches.py", "Branch._apply_sync_update"): (3, "Branch.sync"), + ("apply", "change", "models/branches.py", "Branch._handle_sync_delete"): (1, "Branch.sync"), + ("apply", "", "models/changes.py", "ObjectChange"): (1, "apply.alters_data, no call"), + ("undo", "", "models/changes.py", "ObjectChange"): (1, "undo.alters_data, no call"), + ("deserialize_object", "from utilities.serialization", "models/changes.py", ""): (1, "the import"), + ("update_object", "from netbox_branching.utilities", "models/changes.py", ""): (1, "the import"), + ("deserialize_object", "hasattr(model)", "models/changes.py", "ObjectChange.apply"): (1, "each apply above"), + ("deserialize_object", "model", "models/changes.py", "ObjectChange.apply"): (1, "each apply above"), + ("deserialize_object", "", "models/changes.py", "ObjectChange.apply"): (1, "each apply above"), + ("update_object", "", "models/changes.py", "ObjectChange.apply"): (1, "each apply above"), + ("deserialize_object", "", "models/changes.py", "ObjectChange.undo"): (1, "each undo above"), + ("update_object", "", "models/changes.py", "ObjectChange.undo"): (1, "each undo above"), +} +# Each reference to a name that starts a replay: a merge strategy, a sync helper, or a wrapped method. +REPLAY_ENTRIES = { + ("merge", "strategy_class()", "models/branches.py", "Branch.merge"): (1, "inside the wrapped Branch.merge"), + ("revert", "strategy_class()", "models/branches.py", "Branch.revert"): (1, "inside the wrapped Branch.revert"), + ("_apply_sync_update", "self", "models/branches.py", "Branch.sync"): (1, "inside the wrapped Branch.sync"), + ("_handle_sync_delete", "self", "models/branches.py", "Branch.sync"): (1, "inside the wrapped Branch.sync"), + ("merge", "branch", "jobs.py", "MergeBranchJob.run"): (1, "the wrapped Branch.merge"), + ("revert", "branch", "jobs.py", "RevertBranchJob.run"): (1, "the wrapped Branch.revert"), + ("merge", "", "models/branches.py", "Branch"): (1, "merge.alters_data, no call"), + ("revert", "", "models/branches.py", "Branch"): (1, "revert.alters_data, no call"), + ("revert", "", "models/changes.py", "ObjectChange.migrate"): (1, "its revert flag, not a method"), +} +# The calls that name an attribute by a string. +NAMING_CALLS = frozenset({"getattr", "setattr", "hasattr"}) +# The nodes that open a scope. +SCOPES = (ast.FunctionDef, ast.AsyncFunctionDef, ast.ClassDef, ast.Lambda) +# The fingerprint of each scope that the allow-lists name, in netbox-branching 1.2.1; see fingerprint. +REVIEWED_FINGERPRINTS = { + ("jobs.py", "MergeBranchJob.run"): "d039282f2832f890", + ("jobs.py", "RevertBranchJob.run"): "e0be2037fbdbaca7", + ("merge_strategies/iterative.py", "IterativeMergeStrategy.merge"): "42aa32ddbda92a55", + ("merge_strategies/iterative.py", "IterativeMergeStrategy.revert"): "882b00bb19bab80d", + ("merge_strategies/squash.py", "SquashMergeStrategy.merge"): "c57d4abe51355d00", + ("merge_strategies/squash.py", "SquashMergeStrategy.revert"): "257f571c7d9201bb", + ("models/branches.py", "Branch"): "5c1032613ea2f2a4", + ("models/branches.py", "Branch._apply_sync_update"): "7e235e9385c13ac4", + ("models/branches.py", "Branch._handle_sync_delete"): "ff0d26c6de4ae7df", + ("models/branches.py", "Branch.merge"): "3e0ad757b725963d", + ("models/branches.py", "Branch.revert"): "c736016e1ddc2011", + ("models/branches.py", "Branch.sync"): "395aa970f1eeb7fd", + ("models/changes.py", ""): "6ed9b80466510e19", + ("models/changes.py", "ObjectChange"): "b449c0c1b1db6de6", + ("models/changes.py", "ObjectChange.apply"): "6b9f955fbd2370cf", + ("models/changes.py", "ObjectChange.migrate"): "c1927873bbc41a6d", + ("models/changes.py", "ObjectChange.undo"): "24e38af98d10d67c", +} +# The netbox-branching release whose replay paths the allow-lists and the fingerprints record. +REVIEWED_NETBOX_BRANCHING = "1.2.1" +RE_REVIEW_RELEASE = ( + "The installed netbox-branching is not the release whose replay paths were reviewed. Re-review its replay paths: " + "run the scan, read every new reference, every changed fingerprint and any dynamic dispatch, such as getattr with " + "a variable name. Then update REVIEWED_NETBOX_BRANCHING, the fingerprints and the allow-lists together." +) +# The message of a changed reviewed scope. +RE_REVIEW = ( + "A reviewed scope of netbox-branching changed. A change in it can defer a replay past the wrapper without a new " + "reference. Read the scope again, then update its fingerprint and the allow-lists together." +) + + +def _references(node, names): + """Yield ``(name, receiver)`` for each reference that *node* itself makes to one of *names*.""" + if isinstance(node, ast.Attribute) and node.attr in names: + yield node.attr, ast.unparse(node.value) + elif isinstance(node, ast.Name) and node.id in names: + yield node.id, "" + elif isinstance(node, (ast.Import, ast.ImportFrom)): + source = f"from {'.' * node.level}{node.module or ''}" if isinstance(node, ast.ImportFrom) else "import" + for alias in node.names: + if (name := alias.name.rsplit(".", 1)[-1]) in names: + yield name, source + elif isinstance(node, ast.Call) and getattr(node.func, "id", None) in NAMING_CALLS and len(node.args) > 1: + attribute = node.args[1] + if isinstance(attribute, ast.Constant) and attribute.value in names: + yield attribute.value, f"{node.func.id}({ast.unparse(node.args[0])})" + + +def _scoped_nodes(node, scope=()): + """Yield ``(scope, node)`` for each node under *node*; a function, a class and a lambda open a scope.""" + for child in ast.iter_child_nodes(node): + inner = (*scope, getattr(child, "name", "")) if isinstance(child, SCOPES) else scope + yield inner, child + yield from _scoped_nodes(child, inner) + + +def _modules(package): + """Yield ``(module, tree)`` for each module of *package*, its tests excluded.""" + for path in sorted(package.rglob("*.py")): + module = path.relative_to(package).as_posix() + if not module.startswith("tests/"): + yield module, ast.parse(path.read_text(encoding="utf-8"), filename=str(path)) + + +def replay_references(package, names): + """Count each reference to one of *names* in *package*, its tests excluded, by ``(name, receiver, module, scope)``. + + A reference is an attribute, a name, an imported name, or the string of a getattr, setattr or hasattr call. The + receiver is the object of the attribute or the call, the source of the import, or empty. A function, a class and a + lambda each open a scope, so a deferred call is a site of its own. + """ + sites = collections.Counter() + for module, tree in _modules(package): + for scope, node in _scoped_nodes(tree): + for name, receiver in _references(node, names): + sites[(name, receiver, module, ".".join(scope) or "")] += 1 + return sites + + +def fingerprint(node): + """Hash the source of a function, or of the statements of a class or module outside its nested scopes.""" + own = ( + [node] + if isinstance(node, ast.FunctionDef | ast.AsyncFunctionDef) + else [s for s in node.body if not isinstance(s, SCOPES)] + ) + return hashlib.sha256(ast.unparse(ast.Module(own, [])).encode()).hexdigest()[:16] + + +def scope_fingerprints(package, scopes): + """Return the fingerprint of each ``(module, scope)`` of *scopes* that *package* has, scopes named as above.""" + found = {} + for module, tree in _modules(package): + scope_nodes = [(("",), tree)] + [(s, n) for s, n in _scoped_nodes(tree) if isinstance(n, SCOPES)] + for scope, node in scope_nodes: + if (key := (module, ".".join(scope))) in scopes: + found[key] = fingerprint(node) + return found + + +def reviewed_scopes(): + """Return each ``(module, scope)`` that the allow-lists name.""" + return {(module, scope) for _, _, module, scope in (*REPLAY_SAVES, *REPLAY_ENTRIES)} + + +def reviewed(allowed): + """Return the reviewed count of each site of *allowed*.""" + return collections.Counter({site: count for site, (count, _) in allowed.items()}) @skipUnless(BRANCHING_INSTALLED, BRANCHING_SKIP_REASON) -class BranchTestCase(TransactionTestCase): - """Provision real branches. Each branch is removed when its test ends, because ``--reuse-db`` keeps schemas.""" +class ReplayCallSiteContractTest(SimpleTestCase): + """The installed netbox-branching is the reviewed release; the scan and the fingerprints help review the next one.""" - def provision_branch(self, name, user): - """Return a new branch named *name*, provisioned by *user* as netbox-branching's own tests do.""" - from netbox_branching.models import Branch + maxDiff = None + + def test_the_installed_netbox_branching_is_the_reviewed_release(self): + self.assertEqual(branching.installed_version(), REVIEWED_NETBOX_BRANCHING, RE_REVIEW_RELEASE) + + def assert_reviewed(self, allowed): + """Assert that the installed netbox-branching makes the references of *allowed*, each as often, and no other.""" + import netbox_branching + + found = replay_references(Path(netbox_branching.__file__).parent, {name for name, *_ in allowed}) + self.assertDictEqual(dict(found), dict(reviewed(allowed))) + + def test_each_reference_that_saves_a_replayed_object_is_reviewed(self): + self.assert_reviewed(REPLAY_SAVES) + + def test_each_replay_starts_in_a_wrapped_method(self): + self.assert_reviewed(REPLAY_ENTRIES) + + def test_each_reviewed_scope_is_unchanged(self): + import netbox_branching + + found = scope_fingerprints(Path(netbox_branching.__file__).parent, reviewed_scopes()) + + self.assertEqual(set(REVIEWED_FINGERPRINTS), reviewed_scopes()) + self.assertDictEqual(found, REVIEWED_FINGERPRINTS, RE_REVIEW) + + +class ReplayCallSiteScanTest(SimpleTestCase): + """The scan finds a replay reference wherever a release adds one.""" - branch = Branch(name=name) - branch.save(provision=False) - self.addCleanup(remove_branch, branch) - branch.provision(user=user) - # provision() writes the status with a queryset update, which the instance does not see. - branch.refresh_from_db() - return branch + def scan(self, source, names): + """Return the sites that the scan finds in a package whose module ``added.py`` holds *source*.""" + with tempfile.TemporaryDirectory() as directory: + package = Path(directory) + (package / "added.py").write_text(source, encoding="utf-8") + (package / "tests").mkdir() + (package / "tests" / "test_added.py").write_text("change.apply(None)\n", encoding="utf-8") + return replay_references(package, names) + + def test_the_scan_reports_each_call_with_its_receiver_and_scope(self): + source = ( + "class Strategy:\n def merge(self, change):\n change.apply(self)\n" + "def helper(instance, data):\n def nested():\n update_object(instance, data, using=None)\n" + "change.undo(None)\n" + ) + + self.assertEqual( + self.scan(source, {"apply", "undo", "update_object"}), + { + ("apply", "change", "added.py", "Strategy.merge"): 1, + ("update_object", "", "added.py", "helper.nested"): 1, + ("undo", "change", "added.py", ""): 1, + }, + ) + + def test_an_indirect_reference_is_reported(self): + source = ( + "def indirect(change, branch):\n replay = change.apply\n replay(branch)\n" + "def by_name(change, branch):\n getattr(change, 'apply')(branch)\n" + "from .utilities import update_object as u\n" + ) + + self.assertEqual( + self.scan(source, {"apply", "update_object"}), + { + ("apply", "change", "added.py", "indirect"): 1, + ("apply", "getattr(change)", "added.py", "by_name"): 1, + ("update_object", "from .utilities", "added.py", ""): 1, + }, + ) + + def test_a_second_receiver_and_a_second_call_in_one_scope_are_reported(self): + source = ( + "class Job:\n def run(self, branch, strategy):\n" + " branch.merge(None)\n strategy.merge(None)\n strategy.merge(None)\n" + ) + + self.assertEqual( + self.scan(source, {"merge"}), + {("merge", "branch", "added.py", "Job.run"): 1, ("merge", "strategy", "added.py", "Job.run"): 2}, + ) + + def reviewed_state(self, source): + """Return the references and the fingerprint of ``Strategy.merge`` in a package whose ``added.py`` is *source*.""" + with tempfile.TemporaryDirectory() as directory: + package = Path(directory) + (package / "added.py").write_text(source, encoding="utf-8") + return replay_references(package, {"apply"}), scope_fingerprints(package, {("added.py", "Strategy.merge")}) + + def test_a_deferral_inside_a_reviewed_function_changes_its_fingerprint(self): + reviewed_source = ( + "class Strategy:\n def merge(self, branch, changes):\n" + " for change in changes:\n change.apply(branch)\n" + ) + deferrals = { + "a stored reference called later": ( + "class Strategy:\n def merge(self, branch, changes):\n for change in changes:\n" + " replay = change.apply\n on_commit(lambda: replay(branch))\n" + ), + "a generator": ( + "class Strategy:\n def merge(self, branch, changes):\n" + " return (change.apply(branch) for change in changes)\n" + ), + } + references, fingerprints = self.reviewed_state(reviewed_source) + for deferral, source in deferrals.items(): + with self.subTest(deferral=deferral): + deferred_references, deferred_fingerprints = self.reviewed_state(source) + + self.assertEqual(deferred_references, references) + self.assertNotEqual(deferred_fingerprints, fingerprints) + + def test_a_lambda_is_a_scope_of_its_own(self): + source = "class Branch:\n def merge(self, strategy):\n on_commit(lambda: strategy.merge(self))\n" + + self.assertEqual(self.scan(source, {"merge"}), {("merge", "strategy", "added.py", "Branch.merge."): 1}) class BranchProvisioningTest(BranchTestCase): @@ -113,11 +391,9 @@ def test_a_branch_provisions_ready(self): self.assertTrue(schema_exists(branch.schema_name)) def test_a_write_in_the_active_branch_lands_in_the_branch_schema_only(self): - from netbox_branching.utilities import activate_branch - branch = self.provision_branch("Write", self.user) - with activate_branch(branch): + with activate(branch): self.assertEqual(router.db_for_write(Interface), branch.connection_name) Interface.objects.create(device=self.device, name="branch-only", type="1000base-t") self.assertTrue(Interface.objects.filter(device=self.device, name="branch-only").exists()) @@ -126,10 +402,8 @@ def test_a_write_in_the_active_branch_lands_in_the_branch_schema_only(self): self.assertFalse(Interface.objects.filter(device=self.device, name="branch-only").exists()) def test_the_teardown_drops_the_schema_and_closes_the_connection(self): - from netbox_branching.utilities import activate_branch - branch = self.provision_branch("Teardown", self.user) - with activate_branch(branch): + with activate(branch): self.assertFalse(Interface.objects.filter(device=self.device).exists()) branch_connection = connections[branch.connection_name] self.assertIsNotNone(branch_connection.connection) diff --git a/netbox_interface_name_rules/tests/test_breakout_mode.py b/netbox_interface_name_rules/tests/test_breakout_mode.py index f108718c..072eff11 100644 --- a/netbox_interface_name_rules/tests/test_breakout_mode.py +++ b/netbox_interface_name_rules/tests/test_breakout_mode.py @@ -54,38 +54,30 @@ from netbox_interface_name_rules.models import InterfaceNameRule from netbox_interface_name_rules.name_template import referenced_variables from netbox_interface_name_rules.rule_selection import _VERSION_COLUMNS -from netbox_interface_name_rules.tests.out_of_band import rename_out_of_band -from netbox_interface_name_rules.tests.test_channelization import ( +from netbox_interface_name_rules.tests.helpers import ( + CHANNELIZED, + FLAT, PARENT_TYPE, PLUGIN_LOGGER, REQUIRES_CHANNELIZATION, + TEST_PASSWORD, ChannelizationTestCase, - _build_device, + build_device, + plain_module_type, ) +from netbox_interface_name_rules.tests.out_of_band import rename_out_of_band from netbox_interface_name_rules.views import RulePreview -FLAT = "flat" -CHANNELIZED = "channelized" - -TEST_PASSWORD = "testpass123" # noqa: S105 - Test credential only. - User = get_user_model() -def _plain_module_type(manufacturer, model, iface_type=PARENT_TYPE): - """Create a ModuleType with a single plain (non-channelized) port template.""" - module_type = ModuleType.objects.create(manufacturer=manufacturer, model=model, part_number=model) - InterfaceTemplate.objects.create(module_type=module_type, name="{module}", type=iface_type) - return module_type - - class BreakoutModeFieldTest(TestCase): """The model exposes the mode and the parent template with the defaults existing rules rely on.""" @classmethod def setUpTestData(cls): - manufacturer, cls.device = _build_device("BrkField") - cls.module_type = _plain_module_type(manufacturer, "BrkField-QSFP") + manufacturer, cls.device = build_device("BrkField") + cls.module_type = plain_module_type(manufacturer, "BrkField-QSFP") def test_a_new_rule_defaults_to_the_flat_topology(self): """Nothing changes for a rule that never mentions the mode — flat is what it always did.""" @@ -138,8 +130,8 @@ class BreakoutModeValidationTest(TestCase): @classmethod def setUpTestData(cls): - manufacturer, cls.device = _build_device("BrkValid") - cls.module_type = _plain_module_type(manufacturer, "BrkValid-QSFP") + manufacturer, cls.device = build_device("BrkValid") + cls.module_type = plain_module_type(manufacturer, "BrkValid-QSFP") def _rule(self, **kwargs): """Return an unsaved module rule with *kwargs* applied over sane defaults.""" @@ -278,8 +270,8 @@ class BreakoutModeExportTest(TestCase): @classmethod def setUpTestData(cls): - manufacturer, cls.device = _build_device("BrkExp") - cls.module_type = _plain_module_type(manufacturer, "BrkExp-QSFP") + manufacturer, cls.device = build_device("BrkExp") + cls.module_type = plain_module_type(manufacturer, "BrkExp-QSFP") cls.channelized = InterfaceNameRule.objects.create( module_type=cls.module_type, name_template="xe-0/0/{bay_position}:{channel}", @@ -314,7 +306,7 @@ def test_yaml_export_names_the_mode(self): def test_yaml_export_of_a_flat_rule_still_names_the_mode(self): """flat is a real value, not an absence — an importer must not have to guess it.""" flat = InterfaceNameRule.objects.create( - module_type=_plain_module_type(self.module_type.manufacturer, "BrkExp-QSFP-FLAT"), + module_type=plain_module_type(self.module_type.manufacturer, "BrkExp-QSFP-FLAT"), name_template="xe-0/0/{bay_position}:{channel}", channel_count=4, ) @@ -333,9 +325,9 @@ def setUpTestData(cls): cls.superuser = User.objects.create_superuser( username="brkimport", password=TEST_PASSWORD, email="brkimport@example.com" ) - manufacturer, cls.device = _build_device("BrkImp") - cls.module_type = _plain_module_type(manufacturer, "BrkImp-QSFP") - cls.target_type = _plain_module_type(manufacturer, "BrkImp-QSFP-TARGET") + manufacturer, cls.device = build_device("BrkImp") + cls.module_type = plain_module_type(manufacturer, "BrkImp-QSFP") + cls.target_type = plain_module_type(manufacturer, "BrkImp-QSFP-TARGET") cls.rule = InterfaceNameRule.objects.create( module_type=cls.module_type, name_template="xe-0/0/{bay_position}:{channel}", @@ -382,9 +374,9 @@ def setUpTestData(cls): cls.superuser = User.objects.create_superuser( username="brkbulk", password=TEST_PASSWORD, email="brkbulk@example.com" ) - manufacturer, cls.device = _build_device("BrkBulk") + manufacturer, cls.device = build_device("BrkBulk") cls.rule = InterfaceNameRule.objects.create( - module_type=_plain_module_type(manufacturer, "BrkBulk-QSFP"), + module_type=plain_module_type(manufacturer, "BrkBulk-QSFP"), name_template="xe-0/0/{bay_position}:{channel}", parent_name_template="et-0/0/{bay_position}", breakout_mode=CHANNELIZED, @@ -463,9 +455,9 @@ def setUpTestData(cls): cls.superuser = User.objects.create_superuser( username="brkdetail", password=TEST_PASSWORD, email="brkdetail@example.com" ) - manufacturer, cls.device = _build_device("BrkDetail") + manufacturer, cls.device = build_device("BrkDetail") cls.rule = InterfaceNameRule.objects.create( - module_type=_plain_module_type(manufacturer, "BrkDetail-QSFP"), + module_type=plain_module_type(manufacturer, "BrkDetail-QSFP"), name_template="xe-0/0/{bay_position}:{channel}", parent_name_template="et-XYZZY/0/{bay_position}", breakout_mode=CHANNELIZED, @@ -498,9 +490,9 @@ class BreakoutModeAPITest(APITestCase): @classmethod def setUpTestData(cls): - manufacturer, cls.device = _build_device("BrkApi") - cls.module_type = _plain_module_type(manufacturer, "BrkApi-QSFP") - cls.other_type = _plain_module_type(manufacturer, "BrkApi-QSFP-2") + manufacturer, cls.device = build_device("BrkApi") + cls.module_type = plain_module_type(manufacturer, "BrkApi-QSFP") + cls.other_type = plain_module_type(manufacturer, "BrkApi-QSFP-2") cls.rule = InterfaceNameRule.objects.create( module_type=cls.module_type, name_template="xe-0/0/{bay_position}:{channel}", @@ -571,16 +563,16 @@ class BreakoutModeGraphQLTest(APITestCase): @classmethod def setUpTestData(cls): - manufacturer, cls.device = _build_device("BrkGql") + manufacturer, cls.device = build_device("BrkGql") cls.channelized = InterfaceNameRule.objects.create( - module_type=_plain_module_type(manufacturer, "BrkGql-QSFP-CH"), + module_type=plain_module_type(manufacturer, "BrkGql-QSFP-CH"), name_template="xe-0/0/{bay_position}:{channel}", parent_name_template="et-0/0/{bay_position}", breakout_mode=CHANNELIZED, channel_count=4, ) cls.flat = InterfaceNameRule.objects.create( - module_type=_plain_module_type(manufacturer, "BrkGql-QSFP-FL"), + module_type=plain_module_type(manufacturer, "BrkGql-QSFP-FL"), name_template="xe-0/0/{bay_position}:{channel}", channel_count=4, ) @@ -630,16 +622,16 @@ class BreakoutModeFilterSetTest(TestCase): @classmethod def setUpTestData(cls): - manufacturer, cls.device = _build_device("BrkFs") + manufacturer, cls.device = build_device("BrkFs") cls.match = InterfaceNameRule.objects.create( - module_type=_plain_module_type(manufacturer, "BrkFs-QSFP-A"), + module_type=plain_module_type(manufacturer, "BrkFs-QSFP-A"), name_template="xe-0/0/{bay_position}:{channel}", parent_name_template="et-XYZZY/0/{bay_position}", breakout_mode=CHANNELIZED, channel_count=4, ) cls.other = InterfaceNameRule.objects.create( - module_type=_plain_module_type(manufacturer, "BrkFs-QSFP-B"), + module_type=plain_module_type(manufacturer, "BrkFs-QSFP-B"), name_template="xe-0/0/{bay_position}:{channel}", channel_count=4, ) @@ -659,8 +651,8 @@ def setUpTestData(cls): cls.superuser = User.objects.create_superuser( username="brkform", password=TEST_PASSWORD, email="brkform@example.com" ) - manufacturer, cls.device = _build_device("BrkForm") - cls.module_type = _plain_module_type(manufacturer, "BrkForm-QSFP") + manufacturer, cls.device = build_device("BrkForm") + cls.module_type = plain_module_type(manufacturer, "BrkForm-QSFP") def setUp(self): """Log in before posting to the rule test view.""" @@ -967,8 +959,8 @@ class BreakoutModeFingerprintTest(TestCase): @classmethod def setUpTestData(cls): - manufacturer, cls.device = _build_device("BrkFp") - cls.module_type = _plain_module_type(manufacturer, "BrkFp-QSFP") + manufacturer, cls.device = build_device("BrkFp") + cls.module_type = plain_module_type(manufacturer, "BrkFp-QSFP") cls.rule = InterfaceNameRule.objects.create( module_type=cls.module_type, name_template="xe-0/0/{bay_position}:{channel}", @@ -1005,9 +997,9 @@ class FlatBreakoutModeTest(ChannelizationTestCase): @classmethod def setUpTestData(cls): - manufacturer, cls.device = _build_device("BrkFlat", ["3", "4"]) - cls.explicit_type = _plain_module_type(manufacturer, "BrkFlat-QSFP") - cls.default_type = _plain_module_type(manufacturer, "BrkFlat-QSFP-DEF") + manufacturer, cls.device = build_device("BrkFlat", ["3", "4"]) + cls.explicit_type = plain_module_type(manufacturer, "BrkFlat-QSFP") + cls.default_type = plain_module_type(manufacturer, "BrkFlat-QSFP-DEF") cls.explicit_rule = InterfaceNameRule.objects.create( module_type=cls.explicit_type, name_template="xe-0/0/{bay_position}:{channel}", @@ -1035,7 +1027,7 @@ def setUpTestData(cls): channel_count=4, channel_start=1, ) - cls.second_channel_fails_type = _plain_module_type(manufacturer, "BrkFlat-DIV") + cls.second_channel_fails_type = plain_module_type(manufacturer, "BrkFlat-DIV") # The second channel divides by zero, so the rule can name no family at all. cls.second_channel_fails_rule = InterfaceNameRule.objects.create( module_type=cls.second_channel_fails_type, @@ -1098,7 +1090,7 @@ class FlatBreakoutClaimGateTest(ChannelizationTestCase): @classmethod def setUpTestData(cls): - manufacturer, cls.device = _build_device( + manufacturer, cls.device = build_device( "BrkGate", ["5"], virtual_chassis=VirtualChassis.objects.create(name="brkgate-vc"), vc_position=2 ) cls.module_type = ModuleType.objects.create( @@ -1224,7 +1216,7 @@ class ExecutionOutcomeCoverageTest(ChannelizationTestCase): @classmethod def setUpTestData(cls): - manufacturer, cls.device = _build_device("BrkExec", ["5"]) + manufacturer, cls.device = build_device("BrkExec", ["5"]) cls.module_type = ModuleType.objects.create( manufacturer=manufacturer, model="BrkExec-ZERO", part_number="BrkExec-ZERO" ) @@ -1493,8 +1485,8 @@ class ChannelizedModeWithoutSupportTest(ChannelizationTestCase): @classmethod def setUpTestData(cls): - manufacturer, cls.device = _build_device("BrkNoSup", ["3"]) - cls.module_type = _plain_module_type(manufacturer, "BrkNoSup-QSFP") + manufacturer, cls.device = build_device("BrkNoSup", ["3"]) + cls.module_type = plain_module_type(manufacturer, "BrkNoSup-QSFP") cls.rule = InterfaceNameRule.objects.create( module_type=cls.module_type, name_template="xe-0/0/{bay_position}:{channel}", diff --git a/netbox_interface_name_rules/tests/test_bulk_families.py b/netbox_interface_name_rules/tests/test_bulk_families.py index 3dd67551..2787045a 100644 --- a/netbox_interface_name_rules/tests/test_bulk_families.py +++ b/netbox_interface_name_rules/tests/test_bulk_families.py @@ -5,7 +5,6 @@ from unittest import skipUnless from unittest.mock import MagicMock, patch -from dcim.choices import InterfaceTypeChoices from dcim.models import ( Device, DeviceRole, @@ -47,14 +46,17 @@ template_names, ) from netbox_interface_name_rules.family.names import COLLISION_REASON +from netbox_interface_name_rules.jobs import rule_job_kwargs from netbox_interface_name_rules.models import InterfaceNameRule from netbox_interface_name_rules.tests.committed_callbacks import run_the_reapply -from netbox_interface_name_rules.tests.helpers import make_job, run_job_logged +from netbox_interface_name_rules.tests.helpers import ( + PLAIN_TYPE, + REQUIRES_CHANNELIZATION, + channelized_module_type, + make_job, + run_job_logged, +) from netbox_interface_name_rules.tests.out_of_band import rename_out_of_band -from netbox_interface_name_rules.tests.test_channelization import _channelized_module_type - -PLAIN_TYPE = InterfaceTypeChoices.TYPE_10GE_SFP_PLUS -REQUIRES_CHANNELIZATION = "requires a NetBox that models channelized interfaces" class BulkTestCase(TestCase): @@ -529,7 +531,7 @@ class BulkApplyChannelizedFamiliesTest(BulkTestCase): @classmethod def setUpTestData(cls): super().setUpTestData() - cls.channelized_type = _channelized_module_type(cls.manufacturer, "BULK-CHANNELIZED") + cls.channelized_type = channelized_module_type(cls.manufacturer, "BULK-CHANNELIZED") def setUp(self): self.rule = self._flat_rule(self.channelized_type, name_template="xe-0/0/{bay_position}:{channel}") @@ -790,7 +792,7 @@ def test_a_family_that_can_take_no_name_is_reported_blocked(self): def test_the_background_job_warns_about_what_it_skipped(self): from netbox_interface_name_rules.jobs import ApplyRuleJob - records = run_job_logged(self, ApplyRuleJob(make_job("BulkJob")), rule_id=self.rule.pk) + records = run_job_logged(self, ApplyRuleJob(make_job("BulkJob")), **rule_job_kwargs(self.rule.pk)) self.assertEqual([record.levelname for record in records], ["INFO", "WARNING"]) self.assertEqual(records[1].args, (4,)) diff --git a/netbox_interface_name_rules/tests/test_change_log.py b/netbox_interface_name_rules/tests/test_change_log.py index ac8dcf32..d064549e 100644 --- a/netbox_interface_name_rules/tests/test_change_log.py +++ b/netbox_interface_name_rules/tests/test_change_log.py @@ -17,7 +17,7 @@ from dcim.models import Interface, InterfaceTemplate, Module, ModuleBay from django.contrib.auth import get_user_model from django.contrib.contenttypes.models import ContentType -from django.db import connection +from django.db import DatabaseError, connection from django.test import TestCase, TransactionTestCase from django.test.utils import CaptureQueriesContext from django.urls import reverse @@ -26,9 +26,10 @@ from netbox_interface_name_rules import engine from netbox_interface_name_rules.choices import BreakoutModeChoices -from netbox_interface_name_rules.jobs import ApplyRuleJob, run_as_job_user +from netbox_interface_name_rules.jobs import ApplyRuleJob, rule_job_kwargs, run_as_job_user from netbox_interface_name_rules.models import InterfaceNameRule from netbox_interface_name_rules.tests.helpers import ( + PLAIN_TYPE, empty_the_webhook_queue, make_device, make_device_type, @@ -42,7 +43,6 @@ ) User = get_user_model() -PLAIN_TYPE = "10gbase-x-sfpp" UPDATE = ObjectChangeActionChoices.ACTION_UPDATE @@ -203,7 +203,9 @@ def test_the_tag_records_the_stored_rule_not_the_caller_s_copy(self): copy = InterfaceNameRule.objects.get(pk=self.rule.pk) InterfaceNameRule.objects.filter(pk=self.rule.pk).update(description="edited after the read") - run_as_job_user(make_job("ChgLogFlag"), lambda: engine._flag_rule_potentially_deprecated(copy)) + run_as_job_user( + make_job("ChgLogFlag"), lambda: engine._flag_rule_potentially_deprecated(copy), branch_schema_id=None + ) change = updates_of(self.rule).get() descriptions = (change.prechange_data["description"], change.postchange_data["description"]) @@ -225,7 +227,7 @@ def setUpTestData(cls): def _run_the_job(self): job = make_job("ChgLogJob", self.user) - ApplyRuleJob.handle(job, rule_id=self.rule.pk) + ApplyRuleJob.handle(job, **rule_job_kwargs(self.rule.pk)) job.refresh_from_db() self.assertEqual(job.status, JobStatusChoices.STATUS_COMPLETED) return job @@ -246,13 +248,50 @@ def test_a_job_without_a_user_fails_before_it_renames_anything(self): """The change log cannot name who made a change, so the job makes none.""" job = Job.objects.create(name="Apply rule (test)", job_id=uuid.uuid4()) - ApplyRuleJob.handle(job, rule_id=self.rule.pk) + ApplyRuleJob.handle(job, **rule_job_kwargs(self.rule.pk)) job.refresh_from_db() self.assertEqual(job.status, JobStatusChoices.STATUS_ERRORED) self.assertIn("has no user", job.error) self.assertEqual(Interface.objects.get(module=self.module).name, "0") + def test_a_job_enqueued_before_the_upgrade_fails_before_it_renames_anything(self): + """Its kwargs name no branch and no write alias, and no compatibility path guesses them.""" + job = make_job("ChgLogJob", self.user) + + ApplyRuleJob.handle(job, rule_id=self.rule.pk) + + job.refresh_from_db() + self.assertEqual(job.status, JobStatusChoices.STATUS_ERRORED) + self.assertIn("arguments: 'branch_schema_id' and 'expected_alias'", job.error) + self.assertEqual(Interface.objects.get(module=self.module).name, "0") + + def test_a_job_of_a_branch_that_cannot_be_activated_fails_before_it_renames_anything(self): + """The branch is gone, or netbox-branching is not installed: the job must not run on main instead.""" + job = make_job("ChgLogJob", self.user) + + ApplyRuleJob.handle( + job, rule_id=self.rule.pk, branch_schema_id="gone0001", expected_alias="schema_branch_gone0001" + ) + + job.refresh_from_db() + alias_error = RuntimeError("The write alias is 'default', but the operation expects 'schema_branch_gone0001'.") + self.assertEqual((job.status, job.error), (JobStatusChoices.STATUS_ERRORED, repr(alias_error))) + self.assertEqual(Interface.objects.get(module=self.module).name, "0") + + def test_a_request_processor_that_fails_to_enter_fails_the_job_before_it_renames_anything(self): + """NetBox only warns and goes on without it, so a job whose branch activation fails would run on main.""" + job = make_job("ChgLogJob", self.user) + + with request_processors(*registry["request_processors"], branch_activation_that_fails): + ApplyRuleJob.handle(job, **rule_job_kwargs(self.rule.pk)) + + job.refresh_from_db() + self.assertEqual( + (job.status, job.error), (JobStatusChoices.STATUS_ERRORED, repr(DatabaseError(ACTIVATION_FAILURE))) + ) + self.assertEqual(Interface.objects.get(module=self.module).name, "0") + class JobEventRuleTest(_ModuleFixture, TestCase): """An event rule that matches a change of a job runs its action, as for a change in the web UI.""" @@ -273,7 +312,7 @@ def test_a_webhook_rule_gets_the_rename_and_the_job_completes(self): # django-rq enqueues the webhook when the transaction commits. with self.captureOnCommitCallbacks(execute=True): - ApplyRuleJob.handle(job, rule_id=self.rule.pk) + ApplyRuleJob.handle(job, **rule_job_kwargs(self.rule.pk)) job.refresh_from_db() self.assertEqual((job.status, job.error), (JobStatusChoices.STATUS_COMPLETED, "")) @@ -310,6 +349,14 @@ def event_tracking_without_finally(request): netbox_context.events_queue.set({}) +ACTIVATION_FAILURE = "the branch could not be activated" + + +def branch_activation_that_fails(request): + """A request processor that cannot enter, as netbox-branching's when its query for the branch fails.""" + raise DatabaseError(ACTIVATION_FAILURE) + + @contextmanager def request_processors(*processors): """Register only *processors* while the block runs.""" @@ -333,7 +380,7 @@ def setUpTestData(cls): def _handle(self, rule): job = make_job("ChgLogCtx") - left_set = run_in_a_fresh_context(ApplyRuleJob.handle, job, rule_id=rule.pk) + left_set = run_in_a_fresh_context(ApplyRuleJob.handle, job, **rule_job_kwargs(rule.pk)) job.refresh_from_db() return job.status, left_set diff --git a/netbox_interface_name_rules/tests/test_channelization.py b/netbox_interface_name_rules/tests/test_channelization.py index 53472162..2b73216b 100644 --- a/netbox_interface_name_rules/tests/test_channelization.py +++ b/netbox_interface_name_rules/tests/test_channelization.py @@ -14,19 +14,10 @@ import os from unittest import skipUnless -from dcim.choices import InterfaceTypeChoices from dcim.models import ( - Device, - DeviceRole, - DeviceType, Interface, InterfaceTemplate, - Manufacturer, - Module, - ModuleBay, - ModuleBayTemplate, ModuleType, - Site, VirtualChassis, ) from django.test import TestCase @@ -43,94 +34,17 @@ ) from netbox_interface_name_rules.family import FamilyStatus, execute_installed_plan_set, plan_installed_families from netbox_interface_name_rules.models import InterfaceNameRule - -# Resolved defensively so this module still imports on NetBox releases without channelization. -CHANNEL_TYPE = getattr(InterfaceTypeChoices, "TYPE_CHANNEL", "channel") -PARENT_TYPE = InterfaceTypeChoices.TYPE_40GE_QSFP_PLUS -PLAIN_TYPE = InterfaceTypeChoices.TYPE_10GE_SFP_PLUS - -REQUIRES_CHANNELIZATION = "requires a NetBox that models channelized interfaces (4.7+)" -PLUGIN_LOGGER = "netbox_interface_name_rules" - - -def _build_device(prefix, bay_positions=(), **device_kwargs): - """Create a manufacturer, a device type with module bays at *bay_positions*, and one device.""" - slug = prefix.lower() - manufacturer = Manufacturer.objects.create(name=f"{prefix}Mfg", slug=f"{slug}-mfg") - device_type = DeviceType.objects.create(manufacturer=manufacturer, model=f"{prefix}-Dev", slug=f"{slug}-dev") - for position in bay_positions: - ModuleBayTemplate.objects.create(device_type=device_type, name=f"Bay {position}", position=position) - role = DeviceRole.objects.create(name=f"{prefix}Role", slug=f"{slug}-role") - site = Site.objects.create(name=f"{prefix}Site", slug=f"{slug}-site") - device = Device.objects.create(name=f"{slug}-sw1", device_type=device_type, role=role, site=site, **device_kwargs) - return manufacturer, device - - -def _channelized_family(module_type, parent_name, child_names, channels=4): - """Add one channelized parent template plus its channel templates to *module_type*. - - *child_names* maps a channel_id to the template name that channel takes. - """ - parent = InterfaceTemplate.objects.create( - module_type=module_type, name=parent_name, type=PARENT_TYPE, channels=channels - ) - for channel_id, name in child_names.items(): - InterfaceTemplate.objects.create( - module_type=module_type, - name=name, - type=CHANNEL_TYPE, - parent=parent, - channel_id=channel_id, - ) - return parent - - -def _channelized_module_type(manufacturer, model, channels=4, child_channel_ids=(1, 2, 3, 4), child_names=None): - """Create a ModuleType whose interface templates form a channelized family. - - The parent template is ``{module}`` with *channels* set; each entry in *child_channel_ids* adds a - channel-type template bound to it. *child_names* maps a channel_id to a template name, defaulting - to the upstream ``:`` convention. - """ - module_type = ModuleType.objects.create(manufacturer=manufacturer, model=model, part_number=model) - names = child_names or {channel_id: f"{{module}}:{channel_id}" for channel_id in child_channel_ids} - _channelized_family( - module_type, "{module}", {channel_id: names[channel_id] for channel_id in child_channel_ids}, channels=channels - ) - return module_type - - -class ChannelizationTestCase(TestCase): - """Install helpers shared by the channelized module-install test cases.""" - - def _install(self, module_type, position, run_rules=True): - """Install a module into the bay at *position*; run the post-commit rename unless told not to. - - ``run_rules=False`` leaves the freshly instantiated (raw-named) family in place so a test can - call the engine directly and assert its return value. - """ - bay = ModuleBay.objects.get(device=self.device, name=f"Bay {position}") - if run_rules: - with self.captureOnCommitCallbacks(execute=True): - module = Module.objects.create(device=self.device, module_bay=bay, module_type=module_type) - else: - module = Module.objects.create(device=self.device, module_bay=bay, module_type=module_type) - return module, bay - - @staticmethod - def _names(module): - """Return the sorted interface names of *module*.""" - return sorted(Interface.objects.filter(module=module).values_list("name", flat=True)) - - @staticmethod - def _parent(module): - """Return the channelized parent interface of *module*.""" - return Interface.objects.get(module=module, channels__isnull=False) - - @staticmethod - def _child(module, channel_id): - """Return the channel subinterface of *module* bound to *channel_id*.""" - return Interface.objects.get(module=module, channel_id=channel_id) +from netbox_interface_name_rules.tests.helpers import ( + CHANNEL_TYPE, + PARENT_TYPE, + PLAIN_TYPE, + PLUGIN_LOGGER, + REQUIRES_CHANNELIZATION, + ChannelizationTestCase, + build_device, + channelized_family, + channelized_module_type, +) class LiteralBaseInstallTest(ChannelizationTestCase): @@ -138,7 +52,7 @@ class LiteralBaseInstallTest(ChannelizationTestCase): @classmethod def setUpTestData(cls): - manufacturer, cls.device = _build_device("LiteralBase", ["3"]) + manufacturer, cls.device = build_device("LiteralBase", ["3"]) cls.module_type = ModuleType.objects.create(manufacturer=manufacturer, model="LiteralBase-SFP") def test_install_and_reapply_preserve_arithmetic_braces_in_base(self): @@ -172,10 +86,10 @@ class ChannelizedSimpleRuleTest(ChannelizationTestCase): @classmethod def setUpTestData(cls): - manufacturer, cls.device = _build_device("ChanSimple", ["3", "f4"]) - cls.module_type = _channelized_module_type(manufacturer, "ChanSimple-QSFP") + manufacturer, cls.device = build_device("ChanSimple", ["3", "f4"]) + cls.module_type = channelized_module_type(manufacturer, "ChanSimple-QSFP") # One child carries a free-form name that shares no prefix with the parent template. - cls.free_form_type = _channelized_module_type( + cls.free_form_type = channelized_module_type( manufacturer, "ChanSimple-QSFP-FF", child_names={1: "{module}:1", 2: "{module}:2", 3: "{module}:3", 4: "mgmt-chan"}, @@ -250,8 +164,8 @@ class ChannelizedCollisionTest(ChannelizationTestCase): @classmethod def setUpTestData(cls): - manufacturer, cls.device = _build_device("ChanCol", ["3"]) - cls.module_type = _channelized_module_type(manufacturer, "ChanCol-QSFP") + manufacturer, cls.device = build_device("ChanCol", ["3"]) + cls.module_type = channelized_module_type(manufacturer, "ChanCol-QSFP") cls.rule = InterfaceNameRule.objects.create(module_type=cls.module_type, name_template="et-0/0/{bay_position}") def _occupy(self, name): @@ -309,12 +223,12 @@ class ChannelizedBreakoutRuleTest(ChannelizationTestCase): @classmethod def setUpTestData(cls): - manufacturer, cls.device = _build_device("ChanBrk", ["3", "b4", "e5", "p6", "m8"]) - cls.module_type = _channelized_module_type(manufacturer, "ChanBrk-QSFP") - cls.base_type = _channelized_module_type(manufacturer, "ChanBrk-QSFP-BASE") - cls.empty_type = _channelized_module_type(manufacturer, "ChanBrk-QSFP-EMPTY", child_channel_ids=()) - cls.partial_type = _channelized_module_type(manufacturer, "ChanBrk-QSFP-PART", child_channel_ids=(1, 2)) - cls.mismatch_type = _channelized_module_type(manufacturer, "ChanBrk-QSFP-MM", channels=8) + manufacturer, cls.device = build_device("ChanBrk", ["3", "b4", "e5", "p6", "m8"]) + cls.module_type = channelized_module_type(manufacturer, "ChanBrk-QSFP") + cls.base_type = channelized_module_type(manufacturer, "ChanBrk-QSFP-BASE") + cls.empty_type = channelized_module_type(manufacturer, "ChanBrk-QSFP-EMPTY", child_channel_ids=()) + cls.partial_type = channelized_module_type(manufacturer, "ChanBrk-QSFP-PART", child_channel_ids=(1, 2)) + cls.mismatch_type = channelized_module_type(manufacturer, "ChanBrk-QSFP-MM", channels=8) for module_type in (cls.module_type, cls.empty_type, cls.partial_type, cls.mismatch_type): InterfaceNameRule.objects.create( module_type=module_type, @@ -403,15 +317,15 @@ class ChannelizedEnumerationTest(ChannelizationTestCase): @classmethod def setUpTestData(cls): - manufacturer, cls.device = _build_device("ChanEnum", ["3", "9"]) - cls.module_type = _channelized_module_type(manufacturer, "ChanEnum-QSFP") + manufacturer, cls.device = build_device("ChanEnum", ["3", "9"]) + cls.module_type = channelized_module_type(manufacturer, "ChanEnum-QSFP") # A standalone interface alongside the family, so "families counted once" stays distinguishable # from "interfaces not counted at all". InterfaceTemplate.objects.create(module_type=cls.module_type, name="{module}-mgmt", type=PLAIN_TYPE) cls.rule = InterfaceNameRule.objects.create(module_type=cls.module_type, name_template="et-{base}") # A second family whose template does not feed the current name back in, so "already correctly # named" is a state the rule can actually reach (an "et-{base}" rule renames on every pass). - cls.stable_type = _channelized_module_type(manufacturer, "ChanEnum-QSFP-STABLE") + cls.stable_type = channelized_module_type(manufacturer, "ChanEnum-QSFP-STABLE") cls.stable_rule = InterfaceNameRule.objects.create( module_type=cls.stable_type, name_template="et-0/0/{bay_position}" ) @@ -501,7 +415,7 @@ class ChannelizedDeviceRuleTest(TestCase): @classmethod def setUpTestData(cls): vc = VirtualChassis.objects.create(name="chanvc-vc") - _, cls.device = _build_device("ChanVC", virtual_chassis=vc, vc_position=1) + _, cls.device = build_device("ChanVC", virtual_chassis=vc, vc_position=1) cls.device_type = cls.device.device_type cls.parent = Interface.objects.create(device=cls.device, name="et0", type=PARENT_TYPE, channels=4, module=None) for channel_id in range(1, 5): @@ -576,8 +490,8 @@ class ChannelizedBreakoutStandaloneTest(ChannelizationTestCase): @classmethod def setUpTestData(cls): - manufacturer, cls.device = _build_device("ChanMixed", ["7"]) - cls.module_type = _channelized_module_type(manufacturer, "ChanMixed-QSFP") + manufacturer, cls.device = build_device("ChanMixed", ["7"]) + cls.module_type = channelized_module_type(manufacturer, "ChanMixed-QSFP") InterfaceTemplate.objects.create(module_type=cls.module_type, name="{module}-mgmt", type=PLAIN_TYPE) cls.rule = InterfaceNameRule.objects.create( module_type=cls.module_type, @@ -622,14 +536,14 @@ class ChannelizedPredictionTest(ChannelizationTestCase): @classmethod def setUpTestData(cls): - manufacturer, cls.device = _build_device("ChanPred", ["3", "5", "8"]) - cls.breakout_type = _channelized_module_type(manufacturer, "ChanPred-QSFP-BRK") - cls.simple_type = _channelized_module_type( + manufacturer, cls.device = build_device("ChanPred", ["3", "5", "8"]) + cls.breakout_type = channelized_module_type(manufacturer, "ChanPred-QSFP-BRK") + cls.simple_type = channelized_module_type( manufacturer, "ChanPred-QSFP-SMP", child_names={1: "{module}:1", 2: "{module}:2", 3: "{module}:3", 4: "mgmt-chan"}, ) - cls.mismatch_type = _channelized_module_type(manufacturer, "ChanPred-QSFP-MM", channels=8) + cls.mismatch_type = channelized_module_type(manufacturer, "ChanPred-QSFP-MM", channels=8) InterfaceNameRule.objects.create( module_type=cls.breakout_type, name_template="xe-0/0/{bay_position}:{channel}", @@ -692,8 +606,8 @@ class ChannelizedBreakoutTemplateErrorTest(ChannelizationTestCase): @classmethod def setUpTestData(cls): - manufacturer, cls.device = _build_device("ChanErr", ["4"]) - cls.module_type = _channelized_module_type(manufacturer, "ChanErr-QSFP") + manufacturer, cls.device = build_device("ChanErr", ["4"]) + cls.module_type = channelized_module_type(manufacturer, "ChanErr-QSFP") InterfaceTemplate.objects.create( module_type=cls.module_type, name="mgmt-{module}", @@ -731,18 +645,18 @@ class ChannelizedSuffixRecoveryTest(ChannelizationTestCase): @classmethod def setUpTestData(cls): - manufacturer, cls.device = _build_device("ChanRec", ["1", "2"]) + manufacturer, cls.device = build_device("ChanRec", ["1", "2"]) # Two families on one module type, each with its own suffix convention for the same channel_id. cls.two_family_type = ModuleType.objects.create( manufacturer=manufacturer, model="ChanRec-QSFP-2F", part_number="ChanRec-QSFP-2F" ) - _channelized_family( + channelized_family( cls.two_family_type, "{module}a", {channel_id: f"{{module}}a:{channel_id}" for channel_id in range(1, 5)} ) - _channelized_family( + channelized_family( cls.two_family_type, "{module}b", {channel_id: f"{{module}}b.{channel_id}" for channel_id in range(1, 5)} ) - cls.single_family_type = _channelized_module_type(manufacturer, "ChanRec-QSFP-1F") + cls.single_family_type = channelized_module_type(manufacturer, "ChanRec-QSFP-1F") for module_type in (cls.two_family_type, cls.single_family_type): InterfaceNameRule.objects.create(module_type=module_type, name_template="{base}-x") diff --git a/netbox_interface_name_rules/tests/test_channelized_mode.py b/netbox_interface_name_rules/tests/test_channelized_mode.py index 34db6899..041149de 100644 --- a/netbox_interface_name_rules/tests/test_channelized_mode.py +++ b/netbox_interface_name_rules/tests/test_channelized_mode.py @@ -37,21 +37,19 @@ from netbox_interface_name_rules.family.batch import CHANNELIZED_MODULE_REASON from netbox_interface_name_rules.models import InterfaceNameRule from netbox_interface_name_rules.name_template import TEMPLATE_VARIABLES, NamingContext -from netbox_interface_name_rules.tests.test_breakout_mode import ( +from netbox_interface_name_rules.tests.helpers import ( + CHANNEL_TYPE, CHANNELIZED, FLAT, - TEST_PASSWORD, - _plain_module_type, -) -from netbox_interface_name_rules.tests.test_channelization import ( - CHANNEL_TYPE, PARENT_TYPE, PLAIN_TYPE, PLUGIN_LOGGER, REQUIRES_CHANNELIZATION, + TEST_PASSWORD, ChannelizationTestCase, - _build_device, - _channelized_module_type, + build_device, + channelized_module_type, + plain_module_type, ) User = get_user_model() @@ -63,13 +61,13 @@ class ChannelizedModeInstallTest(ChannelizationTestCase): @classmethod def setUpTestData(cls): - manufacturer, cls.device = _build_device("ChanMode", ["3", "4", "5", "6", "7", "8"]) - cls.named_type = _plain_module_type(manufacturer, "ChanMode-QSFP") - cls.bare_type = _plain_module_type(manufacturer, "ChanMode-QSFP-BARE") - cls.offset_type = _plain_module_type(manufacturer, "ChanMode-QSFP-OFF") - cls.vars_type = _plain_module_type(manufacturer, "ChanMode-QSFP-VARS") - cls.base_type = _plain_module_type(manufacturer, "ChanMode-QSFP-BASE") - cls.conventional_type = _plain_module_type(manufacturer, "ChanMode-QSFP-CONVENTIONAL") + manufacturer, cls.device = build_device("ChanMode", ["3", "4", "5", "6", "7", "8"]) + cls.named_type = plain_module_type(manufacturer, "ChanMode-QSFP") + cls.bare_type = plain_module_type(manufacturer, "ChanMode-QSFP-BARE") + cls.offset_type = plain_module_type(manufacturer, "ChanMode-QSFP-OFF") + cls.vars_type = plain_module_type(manufacturer, "ChanMode-QSFP-VARS") + cls.base_type = plain_module_type(manufacturer, "ChanMode-QSFP-BASE") + cls.conventional_type = plain_module_type(manufacturer, "ChanMode-QSFP-CONVENTIONAL") cls.rule = InterfaceNameRule.objects.create( module_type=cls.named_type, name_template="xe-0/0/{bay_position}:{channel}", @@ -228,8 +226,8 @@ class ChannelizedModePreflightTest(ChannelizationTestCase): @classmethod def setUpTestData(cls): - manufacturer, cls.device = _build_device("ChanPre", ["3", "4", "5"]) - cls.module_type = _plain_module_type(manufacturer, "ChanPre-QSFP") + manufacturer, cls.device = build_device("ChanPre", ["3", "4", "5"]) + cls.module_type = plain_module_type(manufacturer, "ChanPre-QSFP") cls.rule = InterfaceNameRule.objects.create( module_type=cls.module_type, name_template="xe-0/0/{bay_position}:{channel}", @@ -326,8 +324,8 @@ class ChannelizedModeFlatFamilyTest(ChannelizationTestCase): @classmethod def setUpTestData(cls): - manufacturer, cls.device = _build_device("ChanFlat", ["3"]) - cls.module_type = _plain_module_type(manufacturer, "ChanFlat-QSFP") + manufacturer, cls.device = build_device("ChanFlat", ["3"]) + cls.module_type = plain_module_type(manufacturer, "ChanFlat-QSFP") cls.rule = InterfaceNameRule.objects.create( module_type=cls.module_type, name_template="xe-0/0/{base}:{channel}", @@ -411,8 +409,8 @@ class ChannelizedModeRetemplatedFlatFamilyTest(ChannelizationTestCase): @classmethod def setUpTestData(cls): - manufacturer, cls.device = _build_device("ChanReTpl", ["3"]) - cls.module_type = _plain_module_type(manufacturer, "ChanReTpl-QSFP") + manufacturer, cls.device = build_device("ChanReTpl", ["3"]) + cls.module_type = plain_module_type(manufacturer, "ChanReTpl-QSFP") cls.rule = InterfaceNameRule.objects.create( module_type=cls.module_type, name_template="xe-0/0/{bay_position}:{channel}", @@ -485,9 +483,9 @@ class ChannelizedModeExistingFamilyTest(ChannelizationTestCase): @classmethod def setUpTestData(cls): - manufacturer, cls.device = _build_device("ChanExist", ["3", "4"]) - cls.channelized_type = _channelized_module_type(manufacturer, "ChanExist-QSFP-CH") - cls.flat_type = _channelized_module_type(manufacturer, "ChanExist-QSFP-FL") + manufacturer, cls.device = build_device("ChanExist", ["3", "4"]) + cls.channelized_type = channelized_module_type(manufacturer, "ChanExist-QSFP-CH") + cls.flat_type = channelized_module_type(manufacturer, "ChanExist-QSFP-FL") InterfaceNameRule.objects.create( module_type=cls.channelized_type, name_template="xe-0/0/{bay_position}:{channel}", @@ -571,8 +569,8 @@ class ChannelizedModeClaimedPlainInterfaceTest(ChannelizationTestCase): @classmethod def setUpTestData(cls): - manufacturer, cls.device = _build_device("ChanPlain", ["3"]) - cls.module_type = _channelized_module_type(manufacturer, "ChanPlain-QSFP") + manufacturer, cls.device = build_device("ChanPlain", ["3"]) + cls.module_type = channelized_module_type(manufacturer, "ChanPlain-QSFP") InterfaceTemplate.objects.create(module_type=cls.module_type, name="mgmt{module}", type=PLAIN_TYPE) cls.rule = InterfaceNameRule.objects.create( module_type=cls.module_type, @@ -607,8 +605,8 @@ def setUpTestData(cls): cls.superuser = User.objects.create_superuser( username="chanmodeview", password=TEST_PASSWORD, email="chanmodeview@example.com" ) - manufacturer, cls.device = _build_device("ChanPrev", ["3"]) - cls.module_type = _plain_module_type(manufacturer, "ChanPrev-QSFP") + manufacturer, cls.device = build_device("ChanPrev", ["3"]) + cls.module_type = plain_module_type(manufacturer, "ChanPrev-QSFP") cls.rule = InterfaceNameRule.objects.create( module_type=cls.module_type, name_template="xe-0/0/{bay_position}:{channel}", @@ -699,11 +697,11 @@ class ChannelizedModePredictionTest(ChannelizationTestCase): @classmethod def setUpTestData(cls): - manufacturer, cls.device = _build_device("ChanPredMode", ["3", "4", "5", "6"]) - cls.named_type = _plain_module_type(manufacturer, "ChanPredMode-QSFP") - cls.bare_type = _plain_module_type(manufacturer, "ChanPredMode-QSFP-BARE") - cls.family_type = _channelized_module_type(manufacturer, "ChanPredMode-QSFP-FAM") - cls.installed_type = _plain_module_type(manufacturer, "ChanPredMode-QSFP-FLAT") + manufacturer, cls.device = build_device("ChanPredMode", ["3", "4", "5", "6"]) + cls.named_type = plain_module_type(manufacturer, "ChanPredMode-QSFP") + cls.bare_type = plain_module_type(manufacturer, "ChanPredMode-QSFP-BARE") + cls.family_type = channelized_module_type(manufacturer, "ChanPredMode-QSFP-FAM") + cls.installed_type = plain_module_type(manufacturer, "ChanPredMode-QSFP-FLAT") cls.installed_rule = InterfaceNameRule.objects.create( module_type=cls.installed_type, name_template="xe-0/0/{bay_position}:{channel}", @@ -798,8 +796,8 @@ class ChannelizedJuniperE2ETest(ChannelizationTestCase): @classmethod def setUpTestData(cls): - manufacturer, cls.device = _build_device("ChanJnpr", ["5"]) - cls.module_type = _plain_module_type(manufacturer, "QSFP-4X10G-LR-CHAN") + manufacturer, cls.device = build_device("ChanJnpr", ["5"]) + cls.module_type = plain_module_type(manufacturer, "QSFP-4X10G-LR-CHAN") cls.rule = InterfaceNameRule.objects.create( module_type=cls.module_type, device_type=cls.device.device_type, diff --git a/netbox_interface_name_rules/tests/test_ci_workflow.py b/netbox_interface_name_rules/tests/test_ci_workflow.py new file mode 100644 index 00000000..246a221e --- /dev/null +++ b/netbox_interface_name_rules/tests/test_ci_workflow.py @@ -0,0 +1,68 @@ +# SPDX-License-Identifier: Apache-2.0 +# Copyright (C) 2025 Marcin Zieba +"""Tests for the coverage wiring of the CI test workflow.""" + +import tomllib +import unittest +from pathlib import Path + +import yaml + +_PROJECT_ROOT = Path(__file__).resolve().parents[2] + + +def _workflow_jobs(): + workflow = yaml.safe_load((_PROJECT_ROOT / ".github" / "workflows" / "test.yaml").read_text(encoding="utf-8")) + return workflow["jobs"] + + +def _steps_using(job, action): + return [step for step in job["steps"] if step.get("uses", "").startswith(f"{action}@")] + + +class CoverageCombineWorkflowTest(unittest.TestCase): + """Every coverage leg must reach the one job that enforces the gate.""" + + def setUp(self): + self.jobs = _workflow_jobs() + self.test_job = self.jobs["test-netbox"] + self.coverage_job = self.jobs["coverage"] + + def _combine_run(self): + return next(step["run"] for step in self.coverage_job["steps"] if step["name"] == "Combine and check coverage") + + def test_the_combine_job_downloads_the_data_of_every_coverage_leg(self): + legs = [cell["coverage"] for cell in self.test_job["strategy"]["matrix"]["include"] if cell.get("coverage")] + downloads = [step["with"]["name"] for step in _steps_using(self.coverage_job, "actions/download-artifact")] + + self.assertEqual(len(legs), len(set(legs)), "two coverage legs would upload under one artifact name") + self.assertEqual(sorted(f"coverage-data-{leg}" for leg in legs), sorted(downloads)) + self.assertEqual(self.coverage_job["needs"], "test-netbox") + + def test_every_coverage_leg_uploads_its_data(self): + uploads = [step["with"] for step in _steps_using(self.test_job, "actions/upload-artifact")] + + self.assertEqual([upload["name"] for upload in uploads], ["coverage-data-${{ matrix.coverage }}"]) + self.assertEqual(uploads[0]["if-no-files-found"], "error") + + def test_the_combine_step_reads_every_downloaded_file(self): + combine = self._combine_run() + + for step in _steps_using(self.coverage_job, "actions/download-artifact"): + with self.subTest(artifact=step["with"]["name"]): + self.assertIn(f"../{step['with']['path']}/.coverage", combine) + + def test_only_the_combine_job_uploads_to_codecov(self): + uploads = {name: len(_steps_using(job, "codecov/codecov-action")) for name, job in self.jobs.items()} + + self.assertEqual({name: count for name, count in uploads.items() if count}, {"coverage": 1}) + + def test_only_the_combined_report_enforces_the_exact_gate(self): + pyproject = tomllib.loads((_PROJECT_ROOT / "pyproject.toml").read_text(encoding="utf-8")) + combine = self._combine_run() + + self.assertEqual(pyproject["tool"]["coverage"]["report"]["fail_under"], 97) + self.assertEqual(pyproject["tool"]["coverage"]["report"]["precision"], 2) + self.assertIn("--cov-fail-under=0", pyproject["tool"]["pytest"]["ini_options"]["addopts"].split()) + self.assertIn("coverage report\n", combine) + self.assertNotIn("--fail-under", combine) diff --git a/netbox_interface_name_rules/tests/test_conversion.py b/netbox_interface_name_rules/tests/test_conversion.py index 7245d2fb..4c94219c 100644 --- a/netbox_interface_name_rules/tests/test_conversion.py +++ b/netbox_interface_name_rules/tests/test_conversion.py @@ -43,35 +43,31 @@ ) from netbox_interface_name_rules.family import FamilyStatus, execute_conversion, plan_module_conversions from netbox_interface_name_rules.family.template_names import BAY_CHAIN_RELATIONS +from netbox_interface_name_rules.jobs import rule_job_kwargs from netbox_interface_name_rules.models import InterfaceNameRule from netbox_interface_name_rules.naming import build_variables from netbox_interface_name_rules.tests.helpers import ( - empty_the_webhook_queue, - make_interface_webhook_rule, - make_job, - queued_webhooks, -) -from netbox_interface_name_rules.tests.out_of_band import rename_out_of_band -from netbox_interface_name_rules.tests.test_breakout_mode import ( + CHANNEL_TYPE, CHANNELIZED, FLAT, - TEST_PASSWORD, - _plain_module_type, -) -from netbox_interface_name_rules.tests.test_channelization import ( - CHANNEL_TYPE, PARENT_TYPE, PLAIN_TYPE, PLUGIN_LOGGER, REQUIRES_CHANNELIZATION, + REQUIRES_NO_CHANNELIZATION, + TEST_PASSWORD, ChannelizationTestCase, - _build_device, + build_device, + empty_the_webhook_queue, + make_interface_webhook_rule, + make_job, + plain_module_type, + queued_webhooks, ) +from netbox_interface_name_rules.tests.out_of_band import rename_out_of_band User = get_user_model() -REQUIRES_NO_CHANNELIZATION = "requires a NetBox that cannot model channelized interfaces (4.6 and older)" - class ConversionTestCase(ChannelizationTestCase): """Installs flat families with a flat rule, then points that rule at the channelized topology.""" @@ -136,8 +132,8 @@ class ConversionVerdictTest(ConversionTestCase): @classmethod def setUpTestData(cls): - manufacturer, cls.device = _build_device("ConvVerdict", ["3"]) - cls.module_type = _plain_module_type(manufacturer, "ConvVerdict-QSFP") + manufacturer, cls.device = build_device("ConvVerdict", ["3"]) + cls.module_type = plain_module_type(manufacturer, "ConvVerdict-QSFP") cls.rule = cls._flat_rule(cls.module_type) def setUp(self): @@ -193,13 +189,8 @@ def test_finding_the_verdicts_converts_nothing(self): """The preflight really performs the conversion to validate it, so the rollback is the feature.""" pks = dict(Interface.objects.filter(module=self.module).values_list("name", "pk")) - with patch( - "netbox_interface_name_rules.family.conversion.transaction.set_rollback", - wraps=transaction.set_rollback, - ) as set_rollback: - self._verdicts() + self._verdicts() - set_rollback.assert_called_once_with(True) self._assert_still_flat(self.module, "3") self.assertEqual(dict(Interface.objects.filter(module=self.module).values_list("name", "pk")), pks) @@ -298,8 +289,8 @@ class ConversionTest(ConversionTestCase): @classmethod def setUpTestData(cls): - manufacturer, cls.device = _build_device("ConvApply", ["3", "4"]) - cls.module_type = _plain_module_type(manufacturer, "ConvApply-QSFP") + manufacturer, cls.device = build_device("ConvApply", ["3", "4"]) + cls.module_type = plain_module_type(manufacturer, "ConvApply-QSFP") cls.rule = cls._flat_rule(cls.module_type) def setUp(self): @@ -454,8 +445,8 @@ class Ch0MetadataSplitTest(ConversionTestCase): @classmethod def setUpTestData(cls): - manufacturer, cls.device = _build_device("ConvMeta", ["3"]) - cls.module_type = _plain_module_type(manufacturer, "ConvMeta-QSFP") + manufacturer, cls.device = build_device("ConvMeta", ["3"]) + cls.module_type = plain_module_type(manufacturer, "ConvMeta-QSFP") cls.rule = cls._flat_rule(cls.module_type) cls.custom_field = CustomField.objects.create(name="conv_note", type="text") cls.custom_field.object_types.set([ObjectType.objects.get_for_model(Interface)]) @@ -550,8 +541,8 @@ class ConversionPreflightTest(ConversionTestCase): @classmethod def setUpTestData(cls): - manufacturer, cls.device = _build_device("ConvBlock", ["3", "4"]) - cls.module_type = _plain_module_type(manufacturer, "ConvBlock-QSFP") + manufacturer, cls.device = build_device("ConvBlock", ["3", "4"]) + cls.module_type = plain_module_type(manufacturer, "ConvBlock-QSFP") cls.rule = cls._flat_rule(cls.module_type) def setUp(self): @@ -747,8 +738,8 @@ def setUpTestData(cls): cls.superuser = User.objects.create_superuser( username="convview", password=TEST_PASSWORD, email="convview@example.com" ) - manufacturer, cls.device = _build_device("ConvView", ["3"]) - cls.module_type = _plain_module_type(manufacturer, "ConvView-QSFP") + manufacturer, cls.device = build_device("ConvView", ["3"]) + cls.module_type = plain_module_type(manufacturer, "ConvView-QSFP") cls.rule = cls._flat_rule(cls.module_type) def setUp(self): @@ -898,8 +889,8 @@ def setUpTestData(cls): cls.superuser = User.objects.create_superuser( username="convlimit", password=TEST_PASSWORD, email="convlimit@example.com" ) - manufacturer, cls.device = _build_device("ConvLimitView", ["3", "4", "5"]) - cls.module_type = _plain_module_type(manufacturer, "ConvLimitView-QSFP") + manufacturer, cls.device = build_device("ConvLimitView", ["3", "4", "5"]) + cls.module_type = plain_module_type(manufacturer, "ConvLimitView-QSFP") cls.rule = cls._flat_rule(cls.module_type) def setUp(self): @@ -1012,8 +1003,8 @@ class ConversionScanCostTest(ConversionTestCase): @classmethod def setUpTestData(cls): - manufacturer, cls.device = _build_device("ConvCost", ["3", "4"]) - cls.module_type = _plain_module_type(manufacturer, "ConvCost-QSFP") + manufacturer, cls.device = build_device("ConvCost", ["3", "4"]) + cls.module_type = plain_module_type(manufacturer, "ConvCost-QSFP") cls.rule = cls._flat_rule(cls.module_type) def setUp(self): @@ -1057,8 +1048,8 @@ class ConversionScanLimitTest(ConversionTestCase): @classmethod def setUpTestData(cls): - manufacturer, cls.device = _build_device("ConvLimit", ["3", "4", "5"]) - cls.module_type = _plain_module_type(manufacturer, "ConvLimit-QSFP") + manufacturer, cls.device = build_device("ConvLimit", ["3", "4", "5"]) + cls.module_type = plain_module_type(manufacturer, "ConvLimit-QSFP") cls.rule = cls._flat_rule(cls.module_type) def setUp(self): @@ -1148,8 +1139,8 @@ def setUpTestData(cls): cls.superuser = User.objects.create_superuser( username="convlog", password=TEST_PASSWORD, email="convlog@example.com" ) - manufacturer, cls.device = _build_device("ConvLog", ["3"]) - cls.module_type = _plain_module_type(manufacturer, "ConvLog-QSFP") + manufacturer, cls.device = build_device("ConvLog", ["3"]) + cls.module_type = plain_module_type(manufacturer, "ConvLog-QSFP") cls.rule = cls._flat_rule(cls.module_type) cls.fhrp_group = FHRPGroup.objects.create(group_id=81, protocol="vrrp2") @@ -1229,8 +1220,8 @@ class ConversionJobTest(ConversionTestCase): @classmethod def setUpTestData(cls): - manufacturer, cls.device = _build_device("ConvJob", ["3", "4"]) - cls.module_type = _plain_module_type(manufacturer, "ConvJob-QSFP") + manufacturer, cls.device = build_device("ConvJob", ["3", "4"]) + cls.module_type = plain_module_type(manufacturer, "ConvJob-QSFP") cls.rule = cls._flat_rule(cls.module_type) cls.operator = User.objects.create_user(username="convjob-operator") @@ -1240,12 +1231,12 @@ def setUp(self): self.other_module, self.other_bay = self._install(self.module_type, "4") self._switch_to_channelized() - def _run_job(self, **kwargs): + def _run_job(self): """Run the conversion job against a real Job row, the way the worker does.""" from netbox_interface_name_rules.jobs import ConvertFlatFamiliesJob job = Job.objects.create(name="Convert flat families (test)", job_id=uuid.uuid4(), user=self.operator) - ConvertFlatFamiliesJob(job).run(rule_id=self.rule.pk, **kwargs) + ConvertFlatFamiliesJob(job).run(**rule_job_kwargs(self.rule.pk)) return job def test_the_job_converts_every_convertible_family_of_the_rule(self): @@ -1285,7 +1276,7 @@ def test_a_missing_rule_is_logged_rather_than_raised(self): rule_id = self.rule.pk self.rule.delete() - ConvertFlatFamiliesJob(job).run(rule_id=rule_id) + ConvertFlatFamiliesJob(job).run(**rule_job_kwargs(rule_id)) self._assert_still_flat(self.module, "3") @@ -1304,8 +1295,8 @@ class ConversionEventTest(ConversionTestCase): @classmethod def setUpTestData(cls): - manufacturer, cls.device = _build_device("ConvEvent", ["3", "4"]) - cls.module_type = _plain_module_type(manufacturer, "ConvEvent-QSFP") + manufacturer, cls.device = build_device("ConvEvent", ["3", "4"]) + cls.module_type = plain_module_type(manufacturer, "ConvEvent-QSFP") cls.rule = cls._flat_rule(cls.module_type) cls.event_rule = make_interface_webhook_rule("ConvEvent") cls.operator = User.objects.create_user(username="convevent-operator", is_superuser=True) @@ -1331,7 +1322,7 @@ def test_the_job_queues_the_events_of_the_converted_family_only(self): # django-rq enqueues a webhook when the transaction commits. with self.captureOnCommitCallbacks(execute=True): - ConvertFlatFamiliesJob.handle(make_job("ConvEvent", self.operator), rule_id=self.rule.pk) + ConvertFlatFamiliesJob.handle(make_job("ConvEvent", self.operator), **rule_job_kwargs(self.rule.pk)) self._assert_only_the_converted_family_sent_events() @@ -1367,8 +1358,8 @@ class FlatToChannelizedJuniperE2ETest(ConversionTestCase): @classmethod def setUpTestData(cls): - manufacturer, cls.device = _build_device("ConvJnpr", ["3"]) - cls.module_type = _plain_module_type(manufacturer, "QSFP-4X10G-LR-CONV") + manufacturer, cls.device = build_device("ConvJnpr", ["3"]) + cls.module_type = plain_module_type(manufacturer, "QSFP-4X10G-LR-CONV") cls.rule = cls._flat_rule(cls.module_type, device_type=cls.device.device_type) cls.tag = Tag.objects.create(name="ConvJnprTag", slug="convjnpr-tag") @@ -1460,8 +1451,8 @@ def setUpTestData(cls): cls.superuser = User.objects.create_superuser( username="convnosup", password=TEST_PASSWORD, email="convnosup@example.com" ) - manufacturer, cls.device = _build_device("ConvNoSup", ["3"]) - cls.module_type = _plain_module_type(manufacturer, "ConvNoSup-QSFP") + manufacturer, cls.device = build_device("ConvNoSup", ["3"]) + cls.module_type = plain_module_type(manufacturer, "ConvNoSup-QSFP") cls.rule = cls._flat_rule(cls.module_type) def setUp(self): diff --git a/netbox_interface_name_rules/tests/test_documentation.py b/netbox_interface_name_rules/tests/test_documentation.py index d36645aa..de00b7e7 100644 --- a/netbox_interface_name_rules/tests/test_documentation.py +++ b/netbox_interface_name_rules/tests/test_documentation.py @@ -22,6 +22,7 @@ from django.core.management.base import CommandError from netbox_interface_name_rules import template_variable_reference +from netbox_interface_name_rules.branching import SUPPORTED_SERIES from netbox_interface_name_rules.models import InterfaceNameRule from netbox_interface_name_rules.name_template import ( TEMPLATE_VARIABLES, @@ -44,6 +45,12 @@ _PROJECT_ROOT = Path(__file__).resolve().parents[2] + +def _normalised_doc(*path): + """Return the text of the project file at *path* with each run of whitespace collapsed to one space.""" + return " ".join(_PROJECT_ROOT.joinpath(*path).read_text(encoding="utf-8").split()) + + _RE2_AUDIT = importlib.import_module("netbox_interface_name_rules.migrations.0014_validate_re2_patterns") _CONVERTER_OFFSET_SCENARIO = re.compile( @@ -318,7 +325,7 @@ class BaseVariableDocumentationTest(unittest.TestCase): """Keep each module-family base path explicit outside the generated reference.""" def test_each_module_family_base_path_is_documented(self): - guide = (_PROJECT_ROOT / "docs" / "template-variables.md").read_text(encoding="utf-8") + guide = _normalised_doc("docs", "template-variables.md") for statement in ( "In a module rule, `{base}` is the raw template name of the interface the rule renames:", @@ -344,7 +351,7 @@ def test_each_module_family_base_path_is_documented(self): "In a device interface rule, `{base}` is the interface's current name.", ): with self.subTest(statement=statement): - self.assertIn(statement, " ".join(guide.split())) + self.assertIn(statement, guide) class PerformanceDocumentationTest(unittest.TestCase): @@ -415,7 +422,7 @@ def test_rule_priority_lists_every_specificity_score(self): ) def test_configuration_states_that_a_save_outside_a_request_has_no_journal_author(self): - guide = " ".join((_PROJECT_ROOT / "docs" / "configuration.md").read_text(encoding="utf-8").split()) + guide = _normalised_doc("docs", "configuration.md") self.assertIn( "The author is the user of the request that saved the change. " @@ -423,8 +430,14 @@ def test_configuration_states_that_a_save_outside_a_request_has_no_journal_autho guide, ) + def test_configuration_puts_a_script_install_in_a_transaction_on_the_interface_write_alias(self): + guide = _normalised_doc("docs", "configuration.md") + + self.assertIn("wrap the install in `transaction.atomic(using=router.db_for_write(Interface))`.", guide) + self.assertNotIn("wrap the install in `transaction.atomic()`", guide) + def test_configuration_states_whom_the_change_log_of_a_job_names(self): - guide = " ".join((_PROJECT_ROOT / "docs" / "configuration.md").read_text(encoding="utf-8").split()) + guide = _normalised_doc("docs", "configuration.md") self.assertIn( "**Run as Background Job** and **Convert as Background Job** record each change as the user who " @@ -432,6 +445,39 @@ def test_configuration_states_whom_the_change_log_of_a_job_names(self): guide, ) + def test_the_guides_name_the_netbox_branching_series_that_the_version_gate_accepts(self): + statements = { + ("docs", "configuration.md"): ( + f"(https://github.com/netboxlabs/netbox-branching) {SUPPORTED_SERIES} on NetBox 4.7." + ), + ("docs", "installation.md"): f"Optional: netbox-branching {SUPPORTED_SERIES}, on NetBox 4.7.", + ("README.md",): f"(netbox-branching {SUPPORTED_SERIES} on NetBox 4.7).", + ("docs", "index.md"): f"(netbox-branching {SUPPORTED_SERIES} on NetBox 4.7).", + } + + for path, statement in statements.items(): + with self.subTest(path="/".join(path)): + self.assertIn(statement, _normalised_doc(*path)) + + def test_configuration_states_that_a_replay_is_not_a_rename_trigger_and_can_rename_a_kept_channel(self): + guide = _normalised_doc("docs", "configuration.md") + + self.assertIn("so the rename triggers do nothing while netbox-branching replays them.", guide) + self.assertIn( + "A merge, a revert or a sync replays the rename of the parent, so NetBox renames the channels again " + "when the replay commits, and the plugin does not act.", + guide, + ) + + def test_upgrade_guide_says_to_let_the_queued_plugin_jobs_finish(self): + guide = _normalised_doc("docs", "installation.md") + + self.assertIn( + "Let the queued **Run as Background Job** and **Convert as Background Job** jobs finish before you " + "upgrade the plugin.", + guide, + ) + def test_transaction_adr_states_unrelated_failure_behavior(self): adr = (_PROJECT_ROOT / "docs" / "adr" / "0005-execute-each-family-in-its-own-transaction.md").read_text( encoding="utf-8" diff --git a/netbox_interface_name_rules/tests/test_installed_families.py b/netbox_interface_name_rules/tests/test_installed_families.py index 6f8887f9..02cef8fe 100644 --- a/netbox_interface_name_rules/tests/test_installed_families.py +++ b/netbox_interface_name_rules/tests/test_installed_families.py @@ -39,14 +39,15 @@ from netbox_interface_name_rules.family.execution import _lock_family from netbox_interface_name_rules.models import InterfaceNameRule from netbox_interface_name_rules.name_template import evaluate_name_template -from netbox_interface_name_rules.tests.helpers import make_placement +from netbox_interface_name_rules.tests.helpers import ( + CHANNEL_TYPE, + PARENT_TYPE, + PLAIN_TYPE, + REQUIRES_CHANNELIZATION, + make_placement, +) from netbox_interface_name_rules.tests.out_of_band import rename_out_of_band -CHANNEL_TYPE = getattr(InterfaceTypeChoices, "TYPE_CHANNEL", "channel") -PARENT_TYPE = InterfaceTypeChoices.TYPE_40GE_QSFP_PLUS -PLAIN_TYPE = InterfaceTypeChoices.TYPE_10GE_SFP_PLUS -REQUIRES_CHANNELIZATION = "requires a NetBox that models channelized interfaces" - class RejectNthInterfaceUpdate: """Fail the *nth* interface UPDATE so earlier successful writes must also roll back.""" diff --git a/netbox_interface_name_rules/tests/test_misc.py b/netbox_interface_name_rules/tests/test_misc.py index 13059db2..68c4633c 100644 --- a/netbox_interface_name_rules/tests/test_misc.py +++ b/netbox_interface_name_rules/tests/test_misc.py @@ -8,6 +8,7 @@ from django.core.exceptions import ValidationError from django.test import TestCase +from netbox_interface_name_rules.jobs import rule_job_kwargs from netbox_interface_name_rules.models import InterfaceNameRule from netbox_interface_name_rules.tests.helpers import make_job, make_unrunnable_rule, run_job_logged @@ -45,19 +46,13 @@ def test_job_meta_name(self): self.assertEqual(ApplyRuleJob.Meta.name, "Apply Interface Name Rule") - def test_job_run_missing_rule_id(self): - """ApplyRuleJob.run logs warning and returns without error when rule_id is missing.""" - from netbox_interface_name_rules.jobs import ApplyRuleJob - - self.assertEqual(levels(run_job_logged(self, ApplyRuleJob(make_job("JobMissingRule")))), ["WARNING"]) - def test_job_run_nonexistent_rule_id(self): """ApplyRuleJob.run logs warning when rule_id doesn't correspond to a rule.""" from netbox_interface_name_rules.jobs import ApplyRuleJob job = ApplyRuleJob(make_job("JobGoneRule")) - self.assertEqual(levels(run_job_logged(self, job, rule_id=999999)), ["WARNING"]) + self.assertEqual(levels(run_job_logged(self, job, **rule_job_kwargs(999999))), ["WARNING"]) # --------------------------------------------------------------------------- @@ -680,13 +675,13 @@ def _make_job(self): def test_job_run_success_logs_info(self): """ApplyRuleJob.run() with valid rule calls apply_rule_to_existing and logs (lines 30-36).""" - self.assertEqual(levels(run_job_logged(self, self._make_job(), rule_id=self.rule.pk)), ["INFO"]) + self.assertEqual(levels(run_job_logged(self, self._make_job(), **rule_job_kwargs(self.rule.pk))), ["INFO"]) def test_job_run_exception_reraises_and_logs(self): """ApplyRuleJob.run() re-raises exception from apply_rule_to_existing (lines 32-34).""" rule = make_unrunnable_rule("JobX") - records = run_job_logged(self, self._make_job(), raises=ValueError, rule_id=rule.pk) + records = run_job_logged(self, self._make_job(), raises=ValueError, **rule_job_kwargs(rule.pk)) self.assertEqual(levels(records), ["ERROR"]) @@ -717,23 +712,19 @@ def test_job_meta_name_names_the_conversion(self): self.assertIn("Convert", ConvertFlatFamiliesJob.Meta.name) - def test_job_run_without_a_rule_id_is_logged(self): - """An enqueue that lost its argument must not fail the worker.""" - self.assertEqual(levels(run_job_logged(self, self._make_job())), ["WARNING"]) - def test_job_run_with_a_deleted_rule_is_logged(self): """A rule deleted between enqueue and execution is a warning, not a traceback.""" - self.assertEqual(levels(run_job_logged(self, self._make_job(), rule_id=999999)), ["WARNING"]) + self.assertEqual(levels(run_job_logged(self, self._make_job(), **rule_job_kwargs(999999))), ["WARNING"]) def test_job_run_reports_the_family_count(self): """The count is the job's whole output, so it is always logged.""" - self.assertEqual(levels(run_job_logged(self, self._make_job(), rule_id=self.rule.pk)), ["INFO"]) + self.assertEqual(levels(run_job_logged(self, self._make_job(), **rule_job_kwargs(self.rule.pk))), ["INFO"]) def test_job_run_reraises_an_unexpected_failure(self): """A failed conversion has to fail the job, or the operator reads it as done.""" rule = make_unrunnable_rule("ConvJobX") - records = run_job_logged(self, self._make_job(), raises=ValueError, rule_id=rule.pk) + records = run_job_logged(self, self._make_job(), raises=ValueError, **rule_job_kwargs(rule.pk)) self.assertEqual(levels(records), ["ERROR"]) diff --git a/netbox_interface_name_rules/tests/test_module_boundaries.py b/netbox_interface_name_rules/tests/test_module_boundaries.py index cf8106ba..c9aa8bac 100644 --- a/netbox_interface_name_rules/tests/test_module_boundaries.py +++ b/netbox_interface_name_rules/tests/test_module_boundaries.py @@ -58,12 +58,40 @@ ): _MIGRATION_BULK_WRITE, } BULK_WRITE_METHODS = frozenset({"bulk_create", "bulk_update"}) -# The one module that opens atomic blocks: each block there keeps its NetBox events only when it commits. +# The one module that uses Django's transaction and connection state: its blocks span every alias of the scope. TRANSACTIONS_MODULE = PACKAGE / "transactions.py" -TRANSACTION_BLOCK_NAMES = frozenset({"atomic", "savepoint", "savepoint_commit", "savepoint_rollback"}) +TRANSACTION_STATE_NAMES = frozenset( + { + "atomic", + "savepoint", + "savepoint_commit", + "savepoint_rollback", + "clean_savepoints", + "set_rollback", + "get_rollback", + "on_commit", + "get_connection", + "mark_for_rollback_on_error", + } +) +CONNECTION_NAMES = frozenset({"connection", "connections", "router"}) +# The django.db names that production code outside transactions.py must not reach: the transaction module too. +DJANGO_DB_NAMES = CONNECTION_NAMES | {"transaction"} +TRANSACTION_MODULE = "django.db.transaction" +# The position of the alias argument of each call that takes one. +ALIAS_ARGUMENT_POSITIONS = {"on_commit": 1, "get_connection": 0} +# The named exceptions of the design (docs/design/netbox-branching.md, "Guards"), by module. +TRANSACTION_STATE_PERMITS = { + "rename_triggers.py": { + "explicit_alias_calls": frozenset({"on_commit", "get_connection"}), + "permitted_imports": frozenset({"transaction"}), + }, + "models.py": {"permitted_imports": frozenset({"router"})}, +} LANGUAGE_MODULE = PACKAGE / "name_template.py" PYPROJECT = PACKAGE.parent / "pyproject.toml" BRANCHING_MODULE = PACKAGE / "branching.py" +TESTS_PACKAGE = f"{PLUGIN_PACKAGE}.tests" def _family_submodules() -> set[str]: @@ -182,15 +210,81 @@ def _bulk_writes(path: pathlib.Path) -> list[str]: ] -def _transaction_blocks(path: pathlib.Path) -> list[str]: - """Return the source of each use of Django's atomic block or savepoint API, as an attribute or an import.""" +def _module_spellings(records, module: str) -> set[str]: + """Return each dotted name that spells *module* in code whose imports are *records*.""" + spellings = set() + for record in records: + if record.level: + continue + if record.name is None and not record.asname: + # `import a.b` binds `a`, so every module below `a` is spelled in full. + if module == record.bound or module.startswith(f"{record.bound}."): + spellings.add(module) + continue + imported = record.module if record.name is None else f"{record.module}.{record.name}" + if module == imported or module.startswith(f"{imported}."): + spellings.add(record.bound + module[len(imported) :]) + return spellings + + +def _imports_from(record, module: str) -> bool: + """Return whether *record* imports *module*, a module below it, or a name from it.""" + return not record.level and (record.module == module or record.module.startswith(f"{module}.")) + + +def _passes_an_alias(call: ast.Call) -> bool: + """Return whether *call*, of a function in ``ALIAS_ARGUMENT_POSITIONS``, names its alias. + + ``None`` names no alias: Django reads it as ``default``. + """ + position = ALIAS_ARGUMENT_POSITIONS[call.func.attr] + aliases = [*call.args[position : position + 1], *(k.value for k in call.keywords if k.arg == "using")] + return any(not (isinstance(alias, ast.Constant) and alias.value is None) for alias in aliases) + + +def _transaction_state_uses( + path: pathlib.Path, explicit_alias_calls=frozenset(), permitted_imports=frozenset() +) -> list[str]: + """Return the source of each use of Django's transaction or connection state, as an attribute or an import. + + Each import of the ``django.db.transaction`` module counts, whatever it imports. A call in + *explicit_alias_calls* that names its alias, and a ``django.db`` name in *permitted_imports*, are + the named exceptions of a module. + """ tree = ast.parse(path.read_text(encoding="utf-8"), filename=str(path)) - imports = {r.statement for r in import_records(ast.walk(tree)) if r.name in TRANSACTION_BLOCK_NAMES} - return [ - ast.unparse(node) + records = import_records(ast.walk(tree)) + transaction_modules = _module_spellings(records, TRANSACTION_MODULE) + db_modules = _module_spellings(records, "django.db") + permitted_calls = { + id(node.func) for node in ast.walk(tree) - if (isinstance(node, ast.Attribute) and node.attr in TRANSACTION_BLOCK_NAMES) or node in imports - ] + if isinstance(node, ast.Call) + and isinstance(node.func, ast.Attribute) + and node.func.attr in explicit_alias_calls + and _passes_an_alias(node) + } + db_names = DJANGO_DB_NAMES - permitted_imports + imports = { + record.statement + for record in records + if not record.level + and ( + (record.module == TRANSACTION_MODULE and record.name in TRANSACTION_STATE_NAMES) + or (record.module == "django.db" and record.name in db_names) + or ("transaction" in db_names and _imports_from(record, TRANSACTION_MODULE)) + ) + } + uses = [] + for node in ast.walk(tree): + if isinstance(node, ast.Attribute) and id(node) not in permitted_calls: + owner = ast.unparse(node.value) + if (node.attr in TRANSACTION_STATE_NAMES and owner in transaction_modules) or ( + node.attr in db_names and owner in db_modules + ): + uses.append(ast.unparse(node)) + elif node in imports: + uses.append(ast.unparse(node)) + return uses def _netbox_branching_imports(path: pathlib.Path) -> list[str]: @@ -206,6 +300,35 @@ def _netbox_branching_imports(path: pathlib.Path) -> list[str]: ] +def _imports_a_test_module(record, package: str) -> bool: + """Return whether *record*, read in *package*, imports a ``test_*`` module or name, or ``*``, from the test package.""" + module = record.absolute(package) + if record.name == "*": + return module == TESTS_PACKAGE or module.startswith(f"{TESTS_PACKAGE}.") + target = module if record.name is None else f"{module}.{record.name}" + inside = target.removeprefix(f"{TESTS_PACKAGE}.") + return inside != target and any(part.startswith("test_") for part in inside.split(".")) + + +def _test_module_imports(path: pathlib.Path, package: str = TESTS_PACKAGE) -> list[str]: + """Return the source of each import statement in *path*, a module of *package*, that imports a test module.""" + records = import_records(ast.walk(ast.parse(path.read_text(encoding="utf-8"), filename=str(path)))) + return [ + ast.unparse(statement) + for statement in import_statements(r for r in records if _imports_a_test_module(r, package)) + ] + + +def _test_package_violations(tests_root: pathlib.Path) -> set[tuple[str, str]]: + """Return ``(path, statement)`` for each import of a test module in any module of the test package at *tests_root*.""" + violations = set() + for path in sorted(tests_root.rglob("*.py")): + relative = path.relative_to(tests_root) + package = ".".join((TESTS_PACKAGE, *relative.parent.parts)) + violations.update((relative.as_posix(), statement) for statement in _test_module_imports(path, package)) + return violations + + def _production_bulk_writes() -> Counter: """Count each bulk write of the production modules by ``(module, call source)``.""" return Counter( @@ -708,49 +831,182 @@ def test_a_dict_or_set_update_is_not_a_bulk_write(self): class TransactionBlockTest(SimpleTestCase): - """Production code opens an atomic block only through ``transactions.atomic_with_events()``. + """Only ``transactions.py`` uses Django's transaction and connection state. - A bare block that rolls back while the work around it continues leaves its NetBox events in the - request's queue, and NetBox then sends them for changes that the database never kept. + A block there opens every alias of the write scope, and keeps its NetBox events only when it + commits. A bare block opens one alias: in a netbox-branching branch its rollback leaves the writes + of the other connection, and its events stay queued for changes that the database never kept. """ - def test_only_the_transactions_module_opens_an_atomic_block(self): + def _write(self, directory, source): + path = pathlib.Path(directory) / "sample.py" + path.write_text(source, encoding="utf-8") + return path + + def test_only_the_transactions_module_uses_transaction_and_connection_state(self): violations = { - (str(path.relative_to(PACKAGE)), block) + (str(path.relative_to(PACKAGE)), use) for path in _production_modules() if path != TRANSACTIONS_MODULE - for block in _transaction_blocks(path) + for use in _transaction_state_uses( + path, **TRANSACTION_STATE_PERMITS.get(str(path.relative_to(PACKAGE)), {}) + ) } self.assertEqual(violations, set()) - def test_the_transactions_module_opens_one_atomic_block(self): - self.assertEqual(_transaction_blocks(TRANSACTIONS_MODULE), ["transaction.atomic"]) + def test_every_named_exception_is_still_used(self): + for name, permit in TRANSACTION_STATE_PERMITS.items(): + for key in permit: + with self.subTest(module=name, exception=key): + path = PACKAGE / name + without = {other: value for other, value in permit.items() if other != key} + self.assertNotEqual( + _transaction_state_uses(path, **without), _transaction_state_uses(path, **permit) + ) + + def test_the_transactions_module_opens_one_atomic_block_per_alias(self): + atomic = [use for use in _transaction_state_uses(TRANSACTIONS_MODULE) if use.endswith(".atomic")] + + self.assertEqual(atomic, ["transaction.atomic"]) def test_the_detector_reports_every_spelling(self): + imported = "from django.db import transaction" spellings = { - "with transaction.atomic():\n pass\n": ["transaction.atomic"], - "@transaction.atomic\ndef write():\n pass\n": ["transaction.atomic"], + f"{imported}\nwith transaction.atomic():\n pass\n": [imported, "transaction.atomic"], + f"{imported}\n@transaction.atomic\ndef write():\n pass\n": [imported, "transaction.atomic"], "from django.db.transaction import atomic\n": ["from django.db.transaction import atomic"], - "from django.db import transaction as tx\ntx.atomic(using='default')\n": ["tx.atomic"], - "sid = transaction.savepoint()\n": ["transaction.savepoint"], + "from django.db import transaction as tx\ntx.atomic(using='default')\n": [ + "from django.db import transaction as tx", + "tx.atomic", + ], + f"{imported}\nsid = transaction.savepoint()\n": [imported, "transaction.savepoint"], + f"{imported}\ntransaction.savepoint_rollback(sid)\n": [imported, "transaction.savepoint_rollback"], + f"{imported}\ntransaction.clean_savepoints()\n": [imported, "transaction.clean_savepoints"], + f"{imported}\ntransaction.set_rollback(True)\n": [imported, "transaction.set_rollback"], + f"{imported}\ntransaction.get_rollback()\n": [imported, "transaction.get_rollback"], + f"{imported}\ntransaction.on_commit(f)\n": [imported, "transaction.on_commit"], + f"{imported}\ntransaction.get_connection()\n": [imported, "transaction.get_connection"], + f"{imported}\nwith transaction.mark_for_rollback_on_error():\n pass\n": [ + imported, + "transaction.mark_for_rollback_on_error", + ], + "import django.db.transaction as tx\ntx.on_commit(f)\n": [ + "import django.db.transaction as tx", + "tx.on_commit", + ], + "from django import db\ndb.transaction.on_commit(f)\n": ["db.transaction.on_commit", "db.transaction"], + "import django.db.models.deletion\ndjango.db.transaction.on_commit(f)\n": [ + "django.db.transaction.on_commit", + "django.db.transaction", + ], + "from django.db.transaction import on_commit as later\n": [ + "from django.db.transaction import on_commit as later" + ], + "from django.db import connection\n": ["from django.db import connection"], + "from django.db import IntegrityError, connections\n": [ + "from django.db import IntegrityError, connections" + ], + "from django.db import router as db_router\n": ["from django.db import router as db_router"], + "import django.db\ndjango.db.connections['default']\n": ["django.db.connections"], + "from django import db\ndb.router.db_for_write(Model)\n": ["db.router"], } with tempfile.TemporaryDirectory() as directory: - path = pathlib.Path(directory) / "sample.py" for source, expected in spellings.items(): with self.subTest(source=source): - path.write_text(source, encoding="utf-8") - self.assertEqual(_transaction_blocks(path), expected) + self.assertEqual(_transaction_state_uses(self._write(directory, source)), expected) - def test_the_helper_and_the_other_transaction_calls_are_not_reported(self): + def test_the_detector_reports_every_import_of_the_transaction_module(self): + spellings = { + "from django.db import transaction\n": ["from django.db import transaction"], + "from django.db import IntegrityError, transaction as tx\n": [ + "from django.db import IntegrityError, transaction as tx" + ], + "import django.db.transaction\n": ["import django.db.transaction"], + "import django.db.transaction as tx\n": ["import django.db.transaction as tx"], + "def f():\n import django.db.transaction\n": ["import django.db.transaction"], + "from django.db.transaction import TransactionManagementError\n": [ + "from django.db.transaction import TransactionManagementError" + ], + "import django.db\nmodule = django.db.transaction\n": ["django.db.transaction"], + "from django import db\nmodule = db.transaction\n": ["db.transaction"], + } with tempfile.TemporaryDirectory() as directory: - path = pathlib.Path(directory) / "sample.py" - path.write_text( - "with atomic_with_events():\n transaction.set_rollback(True)\ntransaction.on_commit(f)\n", - encoding="utf-8", - ) + for source, expected in spellings.items(): + with self.subTest(source=source): + self.assertEqual(_transaction_state_uses(self._write(directory, source)), expected) - self.assertEqual(_transaction_blocks(path), []) + def test_a_similar_module_and_the_transaction_of_another_object_are_not_reported(self): + source = ( + "import django.db.models.deletion\n" + "from django.db import DEFAULT_DB_ALIAS, IntegrityError, models\n" + "from . import transactions\n" + "from .transactions import atomic_with_events\n" + "import transaction_log\n" + "entry = plan.transaction\n" + "entry = transactions.on_commit\n" + ) + with tempfile.TemporaryDirectory() as directory: + self.assertEqual(_transaction_state_uses(self._write(directory, source)), []) + + def test_the_trigger_module_may_import_the_transaction_module(self): + permit = TRANSACTION_STATE_PERMITS["rename_triggers.py"] + cases = { + "from django.db import DEFAULT_DB_ALIAS, transaction\n": [], + "from django.db import connection, transaction\n": ["from django.db import connection, transaction"], + } + with tempfile.TemporaryDirectory() as directory: + for source, expected in cases.items(): + with self.subTest(source=source): + self.assertEqual(_transaction_state_uses(self._write(directory, source), **permit), expected) + + def test_the_helpers_and_other_objects_are_not_reported(self): + source = ( + "from django.db import IntegrityError, models\n" + "from .transactions import atomic_with_events, on_commit, write_scope\n" + "with write_scope(), atomic_with_events() as block:\n" + " block.set_rollback()\n" + "on_commit(f)\n" + "connection = plan.connection\n" + "connection.run_on_commit.append(entry)\n" + "router = NetBoxRouter()\n" + "router.register('rules', View)\n" + "cache.atomic\n" + ) + with tempfile.TemporaryDirectory() as directory: + self.assertEqual(_transaction_state_uses(self._write(directory, source)), []) + + def test_a_call_that_names_its_alias_is_a_named_exception(self): + cases = { + "transaction.on_commit(f, using=alias)\n": [], + "transaction.on_commit(f, alias)\n": [], + "transaction.get_connection(alias)\n": [], + "transaction.get_connection(using=alias)\n": [], + "transaction.on_commit(f)\n": ["transaction.on_commit"], + "transaction.get_connection()\n": ["transaction.get_connection"], + "transaction.on_commit(f, using=None)\n": ["transaction.on_commit"], + "transaction.get_connection(None)\n": ["transaction.get_connection"], + "transaction.set_rollback(True, using=alias)\n": ["transaction.set_rollback"], + } + with tempfile.TemporaryDirectory() as directory: + for body, expected in cases.items(): + with self.subTest(source=body): + path = self._write(directory, "from django.db import transaction\n" + body) + found = _transaction_state_uses(path, **TRANSACTION_STATE_PERMITS["rename_triggers.py"]) + self.assertEqual(found, expected) + + def test_a_permitted_import_is_a_named_exception_of_its_name_only(self): + cases = { + "from django.db import models, router\n": [], + "from django.db import connections, router\n": ["from django.db import connections, router"], + } + with tempfile.TemporaryDirectory() as directory: + for source, expected in cases.items(): + with self.subTest(source=source): + found = _transaction_state_uses( + self._write(directory, source), **TRANSACTION_STATE_PERMITS["models.py"] + ) + self.assertEqual(found, expected) class NetboxBranchingImportTest(SimpleTestCase): @@ -795,6 +1051,99 @@ def test_a_similar_name_a_relative_import_and_the_app_label_are_not_reported(sel self.assertEqual(_netbox_branching_imports(path), []) +class TestModuleImportTest(SimpleTestCase): + """No module of the test package imports a ``test_*`` module: helpers.py, trigger_cases.py and branch_cases.py share. + + A ``test_*`` name imported from a shared module is refused too: pytest would collect a test function there. A ``*`` + import from the test package is refused, because ``__all__`` can name a test module. + """ + + def test_no_module_of_the_test_package_imports_a_test_module(self): + self.assertEqual(_test_package_violations(PACKAGE / "tests"), set()) + + def test_the_detector_reports_every_spelling(self): + spellings = { + "import netbox_interface_name_rules.tests.test_views\n": [ + "import netbox_interface_name_rules.tests.test_views" + ], + "import os, netbox_interface_name_rules.tests.test_views as views\n": [ + "import os, netbox_interface_name_rules.tests.test_views as views" + ], + "import netbox_interface_name_rules.tests.sub.test_views as views\n": [ + "import netbox_interface_name_rules.tests.sub.test_views as views" + ], + "from netbox_interface_name_rules.tests.test_views import ViewTest\n": [ + "from netbox_interface_name_rules.tests.test_views import ViewTest" + ], + "from netbox_interface_name_rules.tests import helpers, test_views\n": [ + "from netbox_interface_name_rules.tests import helpers, test_views" + ], + "from netbox_interface_name_rules.tests.sub import test_views\n": [ + "from netbox_interface_name_rules.tests.sub import test_views" + ], + "from netbox_interface_name_rules.tests.helpers import test_password\n": [ + "from netbox_interface_name_rules.tests.helpers import test_password" + ], + "from .test_views import ViewTest\n": ["from .test_views import ViewTest"], + "from . import test_views\n": ["from . import test_views"], + "from . import *\n": ["from . import *"], + "from .helpers import *\n": ["from .helpers import *"], + "from netbox_interface_name_rules.tests import *\n": ["from netbox_interface_name_rules.tests import *"], + "def f():\n from .test_views import ViewTest\n": ["from .test_views import ViewTest"], + } + with tempfile.TemporaryDirectory() as directory: + path = pathlib.Path(directory) / "sample.py" + for source, expected in spellings.items(): + with self.subTest(source=source): + path.write_text(source, encoding="utf-8") + self.assertEqual(_test_module_imports(path), expected) + + def test_a_relative_import_resolves_against_the_package_of_its_module(self): + with tempfile.TemporaryDirectory() as directory: + path = pathlib.Path(directory) / "sample.py" + path.write_text("from .test_views import ViewTest\nfrom ..test_views import ViewTest\n", encoding="utf-8") + + reported = _test_module_imports(path, f"{TESTS_PACKAGE}.sub") + + self.assertEqual(reported, ["from .test_views import ViewTest", "from ..test_views import ViewTest"]) + + def test_the_guard_reads_every_module_of_the_test_package_and_its_subpackages(self): + with tempfile.TemporaryDirectory() as directory: + root = pathlib.Path(directory) + (root / "sub").mkdir() + (root / "__init__.py").write_text("from .test_views import ViewTest\n", encoding="utf-8") + (root / "sub" / "__init__.py").write_text("from . import test_views\n", encoding="utf-8") + (root / "sub" / "shared.py").write_text("from ..test_views import ViewTest\n", encoding="utf-8") + (root / "helpers.py").write_text("from .sub import shared\n", encoding="utf-8") + + violations = _test_package_violations(root) + + self.assertEqual( + violations, + { + ("__init__.py", "from .test_views import ViewTest"), + ("sub/__init__.py", "from . import test_views"), + ("sub/shared.py", "from ..test_views import ViewTest"), + }, + ) + + def test_a_shared_module_a_similar_name_and_a_dotted_path_string_are_not_reported(self): + source = ( + "from .helpers import PLAIN_TYPE\n" + "from netbox_interface_name_rules.tests import helpers\n" + "from ..engine import test_rule\n" + "import netbox_interface_name_rules.tests_extra.test_views\n" + "from netbox.settings import *\n" + "from netbox_interface_name_rules.tests_extra import *\n" + "MIDDLEWARE = ('netbox_interface_name_rules.tests.test_views._route',)\n" + ) + with tempfile.TemporaryDirectory() as directory: + path = pathlib.Path(directory) / "sample.py" + path.write_text(source, encoding="utf-8") + + self.assertEqual(_test_module_imports(path), []) + + class UnisolatedReverseTest(SimpleTestCase): """A `SimpleTestCase` must not resolve a URL against NetBox's real root URLconf. diff --git a/netbox_interface_name_rules/tests/test_module_move_trigger.py b/netbox_interface_name_rules/tests/test_module_move_trigger.py index 867f7612..3ec3133f 100644 --- a/netbox_interface_name_rules/tests/test_module_move_trigger.py +++ b/netbox_interface_name_rules/tests/test_module_move_trigger.py @@ -14,15 +14,10 @@ import functools import os import re -from contextlib import contextmanager -from typing import NamedTuple from unittest import skipIf, skipUnless -from unittest.mock import patch -from dcim.models import Interface, InterfaceTemplate, Module, ModuleBay, ModuleBayTemplate, Platform, VirtualChassis -from django.contrib.contenttypes.models import ContentType -from django.db import DataError, IntegrityError, connection, transaction -from django.test import TestCase +from dcim.models import Interface, InterfaceTemplate, Module, ModuleBay, ModuleBayTemplate +from django.db import DataError, connection, transaction from django.test.utils import CaptureQueriesContext from django.urls import reverse from extras.choices import JournalEntryKindChoices @@ -30,7 +25,6 @@ from rest_framework import status from utilities.testing import APITestCase -from netbox_interface_name_rules import engine from netbox_interface_name_rules.choices import BreakoutModeChoices from netbox_interface_name_rules.engine import supports_channelization, supports_vc_position_token from netbox_interface_name_rules.family import supports_module_moves @@ -38,33 +32,39 @@ from netbox_interface_name_rules.rename_triggers import PlanRunner from netbox_interface_name_rules.tests.committed_callbacks import run_the_reapply from netbox_interface_name_rules.tests.helpers import ( - make_device, - make_device_type, - make_manufacturer, + PLAIN_TYPE, + REQUIRES_CHANNELIZATION, + REQUIRES_VC_POSITION_TOKEN, + channelized_module_type, make_module_type, - make_placement, - slug_for, ) from netbox_interface_name_rules.tests.out_of_band import rename_out_of_band -from netbox_interface_name_rules.tests.test_channelization import REQUIRES_CHANNELIZATION, _channelized_module_type -from netbox_interface_name_rules.tests.test_rename_triggers import _give_the_next_module_id, _reject_reads_of -from netbox_interface_name_rules.tests.test_vc_drift import REQUIRES_VC_POSITION_TOKEN +from netbox_interface_name_rules.tests.trigger_cases import ( + CHASSIS_RULES, + FLAT, + NO_RULE, + REQUIRES_SUBTREE_MOVES, + TAKEN, + UNAVAILABLE, + ModuleMoveTestCase, + MoveFixture, + fail_the_naming_read, + give_the_next_module_id, + journal, + module_reapplies, + naming_reads, + reapplied, + reject_interface_updates, + reject_reads_of, +) -PLAIN_TYPE = "10gbase-x-sfpp" -BAYS = (("Bay 0", "0"), ("Bay 1", "1"), ("Bay 2", "2"), ("Bay 10", "10")) -REQUIRES_SUBTREE_MOVES = "requires a NetBox that moves a module's nested bays with it (4.7+)" REQUIRES_DEVICE_MOVES = "requires a NetBox that moves a module's interfaces to its new device (4.7+)" REQUIRES_MOVE_RENAMES = "requires a NetBox that renames a moved module's raw interface names (4.7+)" UNCLAIMED = "no single interface template claims" -FLAT = "a flat breakout family is not renamed after a move, a bay edit or a parent module type change" NOT_RENAMED = "the module is not renamed while one of its interfaces is unclaimed" -NO_RULE = "no rule matches the module after the change" ELSEWHERE = "the interface is not on the device of its module" STALE_BAY = "the module bay still has the parent bay it had before its module moved" -TAKEN = "target name is already in use" -UNAVAILABLE = "{vc_position} is not available on this device" WRITE = re.compile(r'\s*(INSERT INTO|UPDATE|DELETE FROM) "(\w+)"') -NAMING_READ = re.compile(r'SELECT .* FROM "dcim_module" .*"dcim_platform"') def _writes(queries): @@ -72,187 +72,6 @@ def _writes(queries): return [match.groups() for query in queries if (match := WRITE.match(query["sql"]))] -def _reject_interface_updates(execute, sql, params, many, context): - if sql.lstrip().startswith('UPDATE "dcim_interface"'): - raise IntegrityError("injected reapply failure") - return execute(sql, params, many, context) - - -def _fail_the_naming_read(execute, sql, params, many, context): - if NAMING_READ.match(sql): - return execute("SELECT 1/0", None, many, context) - return execute(sql, params, many, context) - - -def _journal(instance): - """Return the journal entries on *instance*, oldest first.""" - return list( - JournalEntry.objects.filter( - assigned_object_type=ContentType.objects.get_for_model(instance), assigned_object_id=instance.pk - ).order_by("pk") - ) - - -@contextmanager -def _module_reapplies(): - """Count the module reapplies; each call still runs the real function.""" - with patch.object(engine, "module_rule_outcomes", wraps=engine.module_rule_outcomes) as spy: - yield spy - - -def _reapplied(spy): - """Return the primary key of the module of each reapply that *spy* recorded, sorted.""" - return sorted(call.args[0].pk for call in spy.call_args_list) - - -class ChassisRule(NamedTuple): - """One rule shape that the tests of a chassis change with a module change cover.""" - - model: str - name_template: str - - -# A plain rule, a rule that reads {base}, and a rule with the virtual-chassis position in arithmetic. -CHASSIS_RULES = ( - ChassisRule("Plain", "et-{vc_position}/{slot}/{bay_position}"), - ChassisRule("Base", "p{base}-{vc_position}/{slot}"), - ChassisRule("Arithmetic", "x{{vc_position} * 10 + {slot_num}}/{bay_position}"), -) - - -@contextmanager -def _naming_reads(): - """Record the queries and the result of each subtree naming read; each call still runs the real function.""" - reads = [] - real = engine.read_subtree_naming - - def read(module_pk): - with CaptureQueriesContext(connection) as queries: - naming = real(module_pk) - reads.append((queries.captured_queries, naming)) - return naming - - with patch.object(engine, "read_subtree_naming", read): - yield reads - - -class _MoveFixture: - """Two device types with the same bays, and three devices in two virtual chassis. - - ``device`` and ``peer`` are virtual-chassis positions 1 and 2 of one chassis; ``remote`` has - another device type and platform, at position 5 of another chassis. - """ - - @classmethod - def build(cls, prefix): - """Create the fixture objects and return them as class attributes of *cls*.""" - cls.prefix = prefix - cls.manufacturer = make_manufacturer(prefix) - cls.device_type = make_device_type(cls.manufacturer, prefix) - cls.other_device_type = make_device_type(cls.manufacturer, f"{prefix} Other") - for device_type in (cls.device_type, cls.other_device_type): - for name, position in BAYS: - ModuleBayTemplate.objects.create(device_type=device_type, name=name, position=position) - cls.platform = Platform.objects.create(name=f"{prefix} OS", slug=slug_for(prefix, "os")) - cls.other_platform = Platform.objects.create(name=f"{prefix} Other OS", slug=slug_for(prefix, "other-os")) - placement = make_placement(prefix) - chassis = VirtualChassis.objects.create(name=f"{prefix} VC") - remote_chassis = VirtualChassis.objects.create(name=f"{prefix} Remote VC") - cls.device = cls._device(placement, "01", cls.device_type, cls.platform, chassis, 1) - cls.peer = cls._device(placement, "02", cls.device_type, cls.other_platform, chassis, 2) - cls.remote = cls._device(placement, "03", cls.other_device_type, cls.other_platform, remote_chassis, 5) - - @classmethod - def _device(cls, placement, suffix, device_type, platform, chassis, position): - return make_device( - cls.prefix, - device_type, - placement, - name=slug_for(cls.prefix, suffix), - platform=platform, - virtual_chassis=chassis, - vc_position=position, - ) - - @classmethod - def _module_type(cls, model, *templates): - module_type = make_module_type(cls.manufacturer, model, model=f"{cls.prefix} {model}") - for template in templates: - InterfaceTemplate.objects.create(module_type=module_type, name=template, type=PLAIN_TYPE) - return module_type - - @classmethod - def _card_type(cls, model, bay_position): - """Return a module type that holds one nested bay at *bay_position*, and no interfaces.""" - card_type = make_module_type(cls.manufacturer, model, model=f"{cls.prefix} {model}") - ModuleBayTemplate.objects.create(module_type=card_type, name="Port", position=bay_position) - return card_type - - @staticmethod - def _bay(device, name="Bay 0"): - return ModuleBay.objects.get(device=device, module__isnull=True, name=name) - - @staticmethod - def _names(module): - return sorted(Interface.objects.filter(module=module).values_list("name", flat=True)) - - -class ModuleMoveTestCase(_MoveFixture, TestCase): - """Install and move modules through real saves, with the committed callbacks run.""" - - @classmethod - def setUpTestData(cls): - cls.build(cls.__name__) - - def _install(self, module_type, bay): - with self.captureOnCommitCallbacks(execute=True): - return Module.objects.create(device=bay.device, module_bay=bay, module_type=module_type) - - @staticmethod - def _save_move(module, bay): - module.device = bay.device - module.module_bay = bay - module.save() - - def _move(self, module, bay): - with self.captureOnCommitCallbacks(execute=True): - self._save_move(module, bay) - - def _install_card(self, card_type, bay): - """Install *card_type* in *bay* and return it with its nested bay.""" - card = self._install(card_type, bay) - return card, ModuleBay.objects.get(module=card) - - def _change_the_chassis_position(self, position=3): - self.device.vc_position = position - self.device.save() - - def _leave_the_chassis(self): - self.device.virtual_chassis = None - self.device.vc_position = None - self.device.save() - - def _join_the_chassis(self, chassis): - self.device.virtual_chassis = chassis - self.device.vc_position = 3 - self.device.save() - - def _save_in_one_transaction(self, *saves): - """Run each of *saves* in order in one transaction; return the spy of the module reapplies.""" - with _module_reapplies() as reapplies, self.captureOnCommitCallbacks(execute=True), transaction.atomic(): - for save in saves: - save() - return reapplies - - def _save_with_a_device_change(self, device_change, save, device_first, before=()): - """Run *device_change* and *save* in one transaction, *device_change* first when *device_first*. - - The saves in *before* run first in the same transaction. - """ - ordered = (device_change, save) if device_first else (save, device_change) - return self._save_in_one_transaction(*before, *ordered) - - class ModuleMoveTest(ModuleMoveTestCase): """A moved module gets the names its rule gives at the new position, for every rule shape.""" @@ -271,7 +90,7 @@ def test_a_module_moved_to_another_bay_is_renamed_for_the_new_bay(self): self._move(module, self._bay(self.device, "Bay 1")) self.assertEqual(self._names(module), ["et-1/0/1"]) - self.assertEqual(_journal(module), []) + self.assertEqual(journal(module), []) @skipUnless(supports_module_moves(), REQUIRES_DEVICE_MOVES) def test_a_module_moved_to_another_device_in_the_chassis_is_renamed_for_its_position(self): @@ -300,7 +119,7 @@ def test_a_base_rule_is_renamed_from_the_raw_name_of_the_new_bay(self): @skipUnless(supports_channelization(), REQUIRES_CHANNELIZATION) def test_a_channelized_family_is_renamed_for_the_new_bay(self): - module_type = _channelized_module_type( + module_type = channelized_module_type( self.manufacturer, f"{self.prefix} Channelized", channels=2, child_channel_ids=(1, 2) ) InterfaceNameRule.objects.create( @@ -320,7 +139,7 @@ def test_a_channelized_family_is_renamed_for_the_new_bay(self): @skipUnless(supports_channelization(), REQUIRES_CHANNELIZATION) def test_a_channelized_family_whose_parent_keeps_its_raw_name_is_renamed_for_the_new_bay(self): - module_type = _channelized_module_type( + module_type = channelized_module_type( self.manufacturer, f"{self.prefix} Kept Parent", channels=2, child_channel_ids=(1, 2) ) InterfaceNameRule.objects.create( @@ -339,7 +158,7 @@ def test_a_channelized_family_whose_parent_keeps_its_raw_name_is_renamed_for_the @skipUnless(supports_channelization(), REQUIRES_CHANNELIZATION) def test_an_unclaimed_interface_beside_a_channelized_family_keeps_every_name(self): - module_type = _channelized_module_type( + module_type = channelized_module_type( self.manufacturer, f"{self.prefix} Beside", channels=2, child_channel_ids=(1, 2) ) InterfaceNameRule.objects.create( @@ -359,7 +178,7 @@ def test_an_unclaimed_interface_beside_a_channelized_family_keeps_every_name(sel callback() self.assertEqual(self._names(module), saved) - (entry,) = _journal(module) + (entry,) = journal(module) self.assertIn(f"`operator-name`: {UNCLAIMED}", entry.comments) for name in saved: if name != "operator-name": @@ -378,7 +197,7 @@ def test_a_subinterface_does_not_stop_the_rename(self): self._move(module, self._bay(self.device, "Bay 1")) self.assertEqual(self._names(module), ["et-1/0/0.100", "et-1/0/1"]) - self.assertEqual(_journal(module), []) + self.assertEqual(journal(module), []) @skipUnless(supports_module_moves(), REQUIRES_SUBTREE_MOVES) @@ -435,7 +254,7 @@ def test_after_a_type_change_and_a_move_the_card_is_reapplied_as_a_type_change_a self._save_move(card, self._bay(self.device, "Bay 2")) self.assertEqual((self._names(card), self._names(optic)), (["p0-1"], ["et-1/2/1"])) - (entry,) = _journal(card) + (entry,) = journal(card) self.assertIn(f"`p0-1`: {UNCLAIMED}", entry.comments) def test_a_move_rolled_back_in_a_savepoint_leaves_the_pending_reapply_as_it_was(self): @@ -472,7 +291,7 @@ def test_a_card_under_a_flat_rule_keeps_its_names_and_its_nested_modules_are_ren self._move(card, self._bay(self.remote, "Bay 1")) self.assertEqual((self._names(card), self._names(optic)), (["a-0:0", "a-0:1"], ["et-5/1/1"])) - (entry,) = _journal(card) + (entry,) = journal(card) self.assertEqual(entry.kind, JournalEntryKindChoices.KIND_WARNING) for name in ("a-0:0", "a-0:1"): self.assertIn(f"`{name}`: {FLAT}", entry.comments) @@ -489,16 +308,16 @@ def test_a_nested_module_whose_reapply_fails_is_reported_with_the_outcomes_befor with self.captureOnCommitCallbacks() as callbacks: self._save_move(card, self._bay(self.remote, "Bay 1")) - with connection.execute_wrapper(_reject_interface_updates), self.assertLogs("netbox_interface_name_rules"): + with connection.execute_wrapper(reject_interface_updates), self.assertLogs("netbox_interface_name_rules"): for callback in callbacks: callback() - (entry,) = _journal(card) + (entry,) = journal(card) self.assertEqual(entry.kind, JournalEntryKindChoices.KIND_DANGER) self.assertIn(f"`et-1/0/1` to `et-5/1/1`: {TAKEN}", entry.comments) self.assertIn("injected reapply failure", entry.comments) self.assertEqual((self._names(blocked), self._names(failed)), (["et-1/0/1"], ["et-1/0/2"])) - self.assertEqual((_journal(blocked), _journal(failed)), ([], [])) + self.assertEqual((journal(blocked), journal(failed)), ([], [])) def test_the_subtree_reports_in_one_journal_entry_on_the_moved_module(self): card, port = self._install_card(self.card_type, self._bay(self.device)) @@ -512,10 +331,10 @@ def test_the_subtree_reports_in_one_journal_entry_on_the_moved_module(self): callback() self.assertEqual(self._names(optic), ["et-1/0/1"]) - (entry,) = _journal(card) + (entry,) = journal(card) self.assertEqual(entry.kind, JournalEntryKindChoices.KIND_WARNING) self.assertIn(f"`et-1/0/1` to `et-5/1/1`: {TAKEN}", entry.comments) - self.assertEqual(_journal(optic), []) + self.assertEqual(journal(optic), []) class RuleWinnerMoveTest(ModuleMoveTestCase): @@ -598,7 +417,7 @@ def test_without_a_rule_after_the_move_the_names_the_old_rule_gave_stay_and_are_ self._move(module, self._bay(self.remote, "Bay 1")) self.assertEqual(self._names(module), ["a0", "operator-name"]) - (entry,) = _journal(module) + (entry,) = journal(module) self.assertEqual(entry.kind, JournalEntryKindChoices.KIND_WARNING) self.assertIn(f"`a0`: {NO_RULE}", entry.comments) self.assertNotIn("operator-name", entry.comments) @@ -619,7 +438,7 @@ def test_a_move_from_a_plain_rule_to_a_flat_breakout_rule_builds_the_family(self self._move(module, self._bay(self.remote, "Bay 1")) self.assertEqual(self._names(module), ["b-1:0", "b-1:1"]) - self.assertEqual(_journal(module), []) + self.assertEqual(journal(module), []) @skipUnless(supports_module_moves(), REQUIRES_DEVICE_MOVES) def test_a_move_into_a_flat_rule_builds_no_family_while_an_interface_is_unclaimed(self): @@ -637,13 +456,13 @@ def test_a_move_into_a_flat_rule_builds_no_family_while_an_interface_is_unclaime self._move(module, self._bay(self.remote, "Bay 1")) self.assertEqual(self._names(module), ["a-0", "a-0:1"]) - (entry,) = _journal(module) + (entry,) = journal(module) self.assertIn(f"`a-0`: {NOT_RENAMED}", entry.comments) self.assertIn(f"`a-0:1`: {UNCLAIMED}", entry.comments) @skipUnless(supports_channelization(), REQUIRES_CHANNELIZATION) def test_without_a_rule_after_the_move_the_channels_renamed_with_their_parent_are_reported(self): - module_type = _channelized_module_type( + module_type = channelized_module_type( self.manufacturer, f"{self.prefix} Lockstep", channels=2, child_channel_ids=(1, 2) ) self._rule("et-{bay_position}", module_type=module_type, device_type=self.device_type) @@ -653,7 +472,7 @@ def test_without_a_rule_after_the_move_the_channels_renamed_with_their_parent_ar self._move(module, self._bay(self.remote, "Bay 1")) self.assertEqual(self._names(module), ["et-0", "et-0:1", "et-0:2"]) - (entry,) = _journal(module) + (entry,) = journal(module) for name in ("et-0", "et-0:1", "et-0:2"): self.assertIn(f"`{name}`: {NO_RULE}", entry.comments) @@ -675,7 +494,7 @@ def test_without_a_rule_after_the_move_the_family_a_channelized_rule_built_is_re self._move(module, self._bay(self.remote, "Bay 1")) self.assertEqual(self._names(module), ["et-0", "et-0:1", "et-0:2"]) - (entry,) = _journal(module) + (entry,) = journal(module) for name in ("et-0", "et-0:1", "et-0:2"): self.assertIn(f"`{name}`: {NO_RULE}", entry.comments) @@ -700,7 +519,7 @@ def test_one_template_that_matches_two_interfaces_renames_neither(self): self._move(module, self._bay(self.device, "Bay 1")) self.assertEqual(self._names(module), ["p0-1", "p1-1"]) - (entry,) = _journal(module) + (entry,) = journal(module) self.assertIn(f"`p0-1`: {UNCLAIMED}", entry.comments) self.assertIn(f"`p1-1`: {UNCLAIMED}", entry.comments) @@ -712,7 +531,7 @@ def test_two_templates_that_match_one_interface_rename_nothing_and_are_reported( self._move(module, self._bay(self.device, "Bay 10")) self.assertEqual(self._names(module), ["operator-name", "x10"]) - (entry,) = _journal(module) + (entry,) = journal(module) self.assertIn(f"`x10`: {UNCLAIMED}", entry.comments) self.assertIn(f"`operator-name`: {UNCLAIMED}", entry.comments) @@ -725,7 +544,7 @@ def test_an_interface_of_a_module_type_without_templates_is_reported_after_a_mov self._move(module, self._bay(self.device, "Bay 1")) self.assertEqual(self._names(module), ["operator-name"]) - (entry,) = _journal(module) + (entry,) = journal(module) self.assertIn(f"`operator-name`: {UNCLAIMED}", entry.comments) def test_an_interface_no_template_matches_keeps_its_name_and_is_reported(self): @@ -735,14 +554,14 @@ def test_an_interface_no_template_matches_keeps_its_name_and_is_reported(self): self._move(module, self._bay(self.device, "Bay 1")) self.assertEqual(self._names(module), ["operator-name"]) - (entry,) = _journal(module) + (entry,) = journal(module) self.assertEqual(entry.kind, JournalEntryKindChoices.KIND_WARNING) self.assertIn(f"`operator-name`: {UNCLAIMED}", entry.comments) def test_a_move_writes_only_the_module_the_rename_and_the_search_cache(self): module = self._install(self.plain_type, self._bay(self.device)) - with _naming_reads() as reads, CaptureQueriesContext(connection) as queries: + with naming_reads() as reads, CaptureQueriesContext(connection) as queries: self._move(module, self._bay(self.device, "Bay 1")) self.assertEqual(self._names(module), ["et-1/0/1"]) @@ -761,7 +580,7 @@ def test_a_save_that_moves_nothing_reads_no_naming(self): module = self._install(self.plain_type, self._bay(self.device)) other_type = self._module_type("Other", "{module}") - with _naming_reads() as reads, self.captureOnCommitCallbacks(execute=True): + with naming_reads() as reads, self.captureOnCommitCallbacks(execute=True): module.description = "unrelated edit" module.save() module.module_type = other_type @@ -774,7 +593,7 @@ def test_a_naming_read_that_fails_fails_the_move_with_its_error(self): module = self._install(self.plain_type, self._bay(self.device)) with ( - connection.execute_wrapper(_fail_the_naming_read), + connection.execute_wrapper(fail_the_naming_read), self.assertRaisesMessage(DataError, "division by zero"), transaction.atomic(), ): @@ -812,7 +631,7 @@ def _move_and_reapply(self, module, bay): def _assert_kept_and_reported(self, module, names): self.assertEqual(self._names(module), names) - (entry,) = _journal(module) + (entry,) = journal(module) self.assertEqual(entry.kind, JournalEntryKindChoices.KIND_WARNING) for name in names: self.assertIn(f"`{name}`: {FLAT}", entry.comments) @@ -851,7 +670,7 @@ def test_a_second_move_under_a_simple_rule_does_not_split_the_family(self): self._move(module, self._bay(self.remote, "Bay 2")) self.assertEqual(self._names(module), ["p1:0", "p1:1"]) - entry = _journal(module)[-1] + entry = journal(module)[-1] self.assertIn(f"`p1:0`: {NOT_RENAMED}", entry.comments) self.assertIn(f"`p1:1`: {UNCLAIMED}", entry.comments) @@ -866,7 +685,7 @@ def test_an_install_and_a_move_in_one_transaction_build_the_family_of_a_flat_rul self._save_move(module, self._bay(self.device, "Bay 2")) self.assertEqual(self._names(module), ["f-2:0", "f-2:1"]) - self.assertEqual(_journal(module), []) + self.assertEqual(journal(module), []) @skipUnless(supports_vc_position_token(), REQUIRES_VC_POSITION_TOKEN) @skipUnless(supports_module_moves(), REQUIRES_DEVICE_MOVES) @@ -896,7 +715,7 @@ def setUpTestData(cls): def test_a_move_in_a_rolled_back_savepoint_causes_no_reapply_and_a_later_move_reapplies_once(self): module = self._install(self.plain_type, self._bay(self.device)) - with _module_reapplies() as reapplies, self.captureOnCommitCallbacks(execute=True), transaction.atomic(): + with module_reapplies() as reapplies, self.captureOnCommitCallbacks(execute=True), transaction.atomic(): with self.assertRaises(RuntimeError), transaction.atomic(): self._save_move(module, self._bay(self.device, "Bay 1")) raise RuntimeError("roll back the savepoint") @@ -909,7 +728,7 @@ def test_a_move_in_a_rolled_back_savepoint_causes_no_reapply_and_a_later_move_re def test_a_move_rolled_back_after_an_earlier_move_keeps_the_earlier_reapply(self): module = self._install(self.plain_type, self._bay(self.device)) - with _module_reapplies() as reapplies, self.captureOnCommitCallbacks(execute=True), transaction.atomic(): + with module_reapplies() as reapplies, self.captureOnCommitCallbacks(execute=True), transaction.atomic(): self._save_move(module, self._bay(self.device, "Bay 2")) with self.assertRaises(RuntimeError), transaction.atomic(): self._save_move(module, self._bay(self.device, "Bay 1")) @@ -922,7 +741,7 @@ def test_a_move_rolled_back_after_an_earlier_move_keeps_the_earlier_reapply(self def test_two_moves_in_one_transaction_reapply_once_from_the_state_before_it(self): module = self._install(self.plain_type, self._bay(self.device)) - with _module_reapplies() as reapplies, self.captureOnCommitCallbacks(execute=True), transaction.atomic(): + with module_reapplies() as reapplies, self.captureOnCommitCallbacks(execute=True), transaction.atomic(): self._save_move(module, self._bay(self.device, "Bay 1")) self._save_move(module, self._bay(self.peer, "Bay 2")) @@ -934,7 +753,7 @@ def test_a_move_the_same_transaction_undoes_reapplies_nothing(self): rename_out_of_band(Interface.objects.get(module=returned), "operator-name") moved = self._install(self.plain_type, self._bay(self.device, "Bay 1")) - with _module_reapplies() as reapplies, self.captureOnCommitCallbacks(execute=True), transaction.atomic(): + with module_reapplies() as reapplies, self.captureOnCommitCallbacks(execute=True), transaction.atomic(): self._save_move(returned, self._bay(self.device, "Bay 2")) self._save_move(returned, self._bay(self.device)) self._save_move(moved, self._bay(self.device, "Bay 10")) @@ -946,14 +765,14 @@ def test_a_bay_edited_before_the_move_in_one_transaction_is_renamed_from_the_sta module = self._install(self.plain_type, self._bay(self.device)) bay = self._bay(self.device) - with _module_reapplies() as reapplies, self.captureOnCommitCallbacks(execute=True), transaction.atomic(): + with module_reapplies() as reapplies, self.captureOnCommitCallbacks(execute=True), transaction.atomic(): bay.position = "5" bay.save() self._save_move(module, self._bay(self.device, "Bay 1")) self.assertEqual(reapplies.call_count, 1) self.assertEqual(self._names(module), ["et-1/0/1"]) - self.assertEqual(_journal(module), []) + self.assertEqual(journal(module), []) class ChassisPositionMoveTest(ModuleMoveTestCase): @@ -1023,8 +842,8 @@ def move_all(): reapplies = self._save_with_a_device_change(device_change, move_all, device_first) - self.assertEqual(_reapplied(reapplies), sorted(module.pk for module in (*modules, other))) - self.assertEqual([_journal(module) for module in modules], [[], [], []]) + self.assertEqual(reapplied(reapplies), sorted(module.pk for module in (*modules, other))) + self.assertEqual([journal(module) for module in modules], [[], [], []]) return [self._names(module) for module in modules], self._names(other) def _device_bays(self): @@ -1037,7 +856,7 @@ def _assert_moved_with_a_position_change(self, targets, chassis_first, names): moved, other = self._move_all(targets, self._change_the_chassis_position, chassis_first) self.assertEqual((moved, other), (names, ["et-3/10/10"])) - self.assertEqual(_journal(self.device), []) + self.assertEqual(journal(self.device), []) def test_moves_then_a_chassis_position_change_rename_each_module_once(self): self._assert_moved_with_a_position_change(self._device_bays(), False, [["et-3/5/5"], ["p6-3/6"], ["x37/7"]]) @@ -1068,14 +887,14 @@ def move_all(): reapplies = self._save_in_one_transaction(self._change_the_chassis_position, move_all) self.assertEqual([self._names(module) for module in modules], [["et-3/5/5"], ["p6-3/6"], ["x37/7"]]) - self.assertEqual(_reapplied(reapplies), sorted(module.pk for module in modules)) - self.assertEqual([_journal(module) for module in (*modules, self.device)], [[], [], [], []]) + self.assertEqual(reapplied(reapplies), sorted(module.pk for module in modules)) + self.assertEqual([journal(module) for module in (*modules, self.device)], [[], [], [], []]) def _assert_moved_out_with_a_leave(self, leave_first): moved, other = self._move_all(self._peer_bays(), self._leave_the_chassis, leave_first) self.assertEqual((moved, other), ([["et-2/0/0"], ["p1-2/1"], ["x22/2"]], ["et-1/10/10"])) - (entry,) = _journal(self.device) + (entry,) = journal(self.device) self.assertIn(f"`et-1/10/10`: {UNAVAILABLE}", entry.comments) @skipUnless(supports_module_moves(), REQUIRES_DEVICE_MOVES) @@ -1100,8 +919,8 @@ def _assert_an_undone_move_leaves_the_module_to_the_chassis_position_change(self reapplies = self._save_in_one_transaction(*saves) - self.assertEqual((self._names(module), _reapplied(reapplies)), (["et-3/0/0"], [module.pk])) - self.assertEqual((_journal(module), _journal(self.device)), ([], [])) + self.assertEqual((self._names(module), reapplied(reapplies)), (["et-3/0/0"], [module.pk])) + self.assertEqual((journal(module), journal(self.device)), ([], [])) # Only the module unit passes the naming points of the moves. moves = [point.move for call in reapplies.call_args_list for point in call.kwargs.get("naming_points", ())] self.assertEqual(any(moves), change_at == 1 and supports_module_moves()) @@ -1131,7 +950,7 @@ def _assert_joining_with_a_move_renames_once(self, join_first): ) self.assertEqual((self._names(module), self._names(other)), (["et-3/5/5"], ["et-3/10/10"])) - self.assertEqual(_reapplied(reapplies), sorted((module.pk, other.pk))) + self.assertEqual(reapplied(reapplies), sorted((module.pk, other.pk))) self.assertEqual(JournalEntry.objects.count(), entries) def test_a_move_then_joining_a_chassis_rename_each_module_once(self): @@ -1149,8 +968,8 @@ def _assert_a_raw_name_is_renamed_once_after_a_move_into_a_rule(self, chassis_fi self._change_the_chassis_position, functools.partial(self._save_move, module, port), chassis_first ) - self.assertEqual((self._names(module), _reapplied(reapplies).count(module.pk)), (["et-3/1/1"], 1)) - self.assertEqual((_journal(module), _journal(self.device)), ([], [])) + self.assertEqual((self._names(module), reapplied(reapplies).count(module.pk)), (["et-3/1/1"], 1)) + self.assertEqual((journal(module), journal(self.device)), ([], [])) @skipUnless(supports_vc_position_token(), REQUIRES_VC_POSITION_TOKEN) def test_a_move_into_a_rule_then_a_chassis_position_change_rename_a_raw_name_that_reads_the_position(self): @@ -1169,8 +988,8 @@ def test_a_name_netbox_gives_at_a_move_is_recognised_after_a_later_position_chan functools.partial(self._save_move, module, port), self._change_the_chassis_position ) - self.assertEqual((self._names(module), _reapplied(reapplies).count(module.pk)), (["et-3/1/1"], 1)) - self.assertEqual((_journal(module), _journal(self.device)), ([], [])) + self.assertEqual((self._names(module), reapplied(reapplies).count(module.pk)), (["et-3/1/1"], 1)) + self.assertEqual((journal(module), journal(self.device)), ([], [])) @skipUnless(supports_module_moves(), REQUIRES_MOVE_RENAMES) def test_a_move_off_a_chassis_then_a_join_recognise_the_name_netbox_gave_at_the_move(self): @@ -1185,8 +1004,8 @@ def test_a_move_off_a_chassis_then_a_join_recognise_the_name_netbox_gave_at_the_ functools.partial(self._save_move, module, port), functools.partial(self._join_the_chassis, chassis) ) - self.assertEqual((self._names(module), _reapplied(reapplies).count(module.pk)), (["et-3/1/2"], 1)) - self.assertEqual((_journal(module), _journal(self.device)), ([], [])) + self.assertEqual((self._names(module), reapplied(reapplies).count(module.pk)), (["et-3/1/2"], 1)) + self.assertEqual((journal(module), journal(self.device)), ([], [])) @skipUnless(supports_module_moves(), REQUIRES_MOVE_RENAMES) def test_a_card_moved_to_another_device_then_its_position_change_recognise_the_nested_name(self): @@ -1202,8 +1021,8 @@ def renumber_the_peer(): functools.partial(self._save_move, card, self._bay(self.peer, "Bay 2")), renumber_the_peer ) - self.assertEqual((self._names(optic), _reapplied(reapplies).count(optic.pk)), (["et-4/2/1"], 1)) - self.assertEqual((_journal(card), _journal(optic), _journal(self.peer)), ([], [], [])) + self.assertEqual((self._names(optic), reapplied(reapplies).count(optic.pk)), (["et-4/2/1"], 1)) + self.assertEqual((journal(card), journal(optic), journal(self.peer)), ([], [], [])) @skipUnless(supports_module_moves(), REQUIRES_MOVE_RENAMES) def test_a_move_into_a_bay_whose_position_is_the_token_then_its_position_change_recognise_the_name(self): @@ -1217,8 +1036,8 @@ def renumber_the_peer(): reapplies = self._save_in_one_transaction(functools.partial(self._save_move, module, bay), renumber_the_peer) - self.assertEqual((self._names(module), _reapplied(reapplies).count(module.pk)), (["et-4/7"], 1)) - self.assertEqual((_journal(module), _journal(self.peer)), ([], [])) + self.assertEqual((self._names(module), reapplied(reapplies).count(module.pk)), (["et-4/7"], 1)) + self.assertEqual((journal(module), journal(self.peer)), ([], [])) @skipUnless(supports_module_moves(), REQUIRES_MOVE_RENAMES) def test_two_moves_around_a_chassis_position_change_recognise_the_name_of_the_first_move(self): @@ -1233,8 +1052,8 @@ def test_two_moves_around_a_chassis_position_change_recognise_the_name_of_the_fi functools.partial(self._save_move, module, self._bay(self.device, "Bay 2")), ) - self.assertEqual((self._names(module), _reapplied(reapplies)), (["et-3/2"], [module.pk])) - self.assertEqual((_journal(module), _journal(self.device)), ([], [])) + self.assertEqual((self._names(module), reapplied(reapplies)), (["et-3/2"], [module.pk])) + self.assertEqual((journal(module), journal(self.device)), ([], [])) def _install_a_raw_module_with_a_base_rule(self): """Install a module whose raw name reads the position in Bay 0, then enable a {base} rule for it.""" @@ -1254,8 +1073,8 @@ def test_a_move_out_and_back_around_a_chassis_position_change_recognise_the_name functools.partial(self._save_move, module, self._bay(self.device)), ) - self.assertEqual((self._names(module), _reapplied(reapplies)), (["p3/0"], [module.pk])) - self.assertEqual((_journal(module), _journal(self.device)), ([], [])) + self.assertEqual((self._names(module), reapplied(reapplies)), (["p3/0"], [module.pk])) + self.assertEqual((journal(module), journal(self.device)), ([], [])) def _install_a_raw_optic_in_a_card_with_a_base_rule(self): """Install a card whose port position is the card's bay, a raw optic in it, and a {base} rule for the optic.""" @@ -1276,8 +1095,8 @@ def test_a_card_moved_out_and_back_around_a_chassis_position_change_recognise_th functools.partial(self._save_move, card, self._bay(self.device)), ) - self.assertEqual((self._names(optic), _reapplied(reapplies).count(optic.pk)), (["p3/0"], 1)) - self.assertEqual((_journal(card), _journal(optic), _journal(self.device)), ([], [], [])) + self.assertEqual((self._names(optic), reapplied(reapplies).count(optic.pk)), (["p3/0"], 1)) + self.assertEqual((journal(card), journal(optic), journal(self.device)), ([], [], [])) @skipUnless(supports_module_moves(), REQUIRES_MOVE_RENAMES) def test_a_card_moved_out_and_back_around_a_port_edit_that_is_undone_recognise_the_nested_name(self): @@ -1295,8 +1114,8 @@ def edit_the_port(position): functools.partial(edit_the_port, "0"), ) - self.assertEqual((self._names(optic), _reapplied(reapplies)), (["p1/0"], [optic.pk])) - self.assertEqual((_journal(card), _journal(optic)), ([], [])) + self.assertEqual((self._names(optic), reapplied(reapplies)), (["p1/0"], [optic.pk])) + self.assertEqual((journal(card), journal(optic)), ([], [])) @skipUnless(supports_module_moves(), REQUIRES_MOVE_RENAMES) def test_a_position_change_of_a_device_the_module_left_before_its_return_recognise_the_name_given_there(self): @@ -1313,8 +1132,8 @@ def renumber_the_peer(): functools.partial(self._save_move, module, self._bay(self.device)), ) - self.assertEqual((self._names(module), _reapplied(reapplies)), (["p1/0"], [module.pk])) - self.assertEqual((_journal(module), _journal(self.device), _journal(self.peer)), ([], [], [])) + self.assertEqual((self._names(module), reapplied(reapplies)), (["p1/0"], [module.pk])) + self.assertEqual((journal(module), journal(self.device), journal(self.peer)), ([], [], [])) @skipUnless(supports_module_moves(), REQUIRES_MOVE_RENAMES) def test_a_move_out_an_edit_of_the_new_bay_and_a_move_back_recognise_the_name_of_the_move_out(self): @@ -1332,8 +1151,8 @@ def edit_the_new_bay(): functools.partial(self._save_move, module, self._bay(self.device)), ) - self.assertEqual((self._names(module), _reapplied(reapplies)), (["p1/0"], [module.pk])) - self.assertEqual(_journal(module), []) + self.assertEqual((self._names(module), reapplied(reapplies)), (["p1/0"], [module.pk])) + self.assertEqual(journal(module), []) @skipUnless(supports_module_moves(), REQUIRES_MOVE_RENAMES) def test_a_move_with_a_position_change_undone_ends_as_the_move_alone(self): @@ -1344,7 +1163,7 @@ def test_a_move_with_a_position_change_undone_ends_as_the_move_alone(self): for device_changes in ((), undone): module = self._install(self.adjacent_type, self._bay(self.device)) self._save_in_one_transaction(functools.partial(self._save_move, module, port), *device_changes) - outcomes.append((self._names(module), [entry.comments for entry in _journal(module)])) + outcomes.append((self._names(module), [entry.comments for entry in journal(module)])) module.delete() self.assertEqual(outcomes, [(["et-1/1/1"], []), (["et-1/1/1"], [])]) @@ -1359,10 +1178,10 @@ def _assert_a_collision_is_reported_once(self, chassis_first): chassis_first, ) - self.assertEqual((self._names(module), _reapplied(reapplies)), (["et-1/0/0"], [module.pk])) - (entry,) = _journal(module) + self.assertEqual((self._names(module), reapplied(reapplies)), (["et-1/0/0"], [module.pk])) + (entry,) = journal(module) self.assertEqual(entry.comments.count(f"`et-1/0/0` to `et-3/5/5`: {TAKEN}"), 1) - self.assertEqual(_journal(self.device), []) + self.assertEqual(journal(self.device), []) def test_a_move_then_a_chassis_position_change_report_a_collision_once(self): self._assert_a_collision_is_reported_once(chassis_first=False) @@ -1413,8 +1232,8 @@ def _assert_an_adjacent_token_install_is_named_once(self, chassis_first, bay, na module, other, reapplies = self._install_with_the_chassis_change(chassis_first, self.adjacent_type, bay) self.assertEqual((self._names(module), self._names(other)), (names, ["et-3/10/10"])) - self.assertEqual(_reapplied(reapplies), sorted((module.pk, other.pk, *(card.pk for card in cards)))) - self.assertEqual((_journal(module), _journal(self.device)), ([], [])) + self.assertEqual(reapplied(reapplies), sorted((module.pk, other.pk, *(card.pk for card in cards)))) + self.assertEqual((journal(module), journal(self.device)), ([], [])) @skipUnless(supports_vc_position_token(), REQUIRES_VC_POSITION_TOKEN) def test_an_install_then_a_chassis_position_change_recognise_the_raw_name_at_the_position_of_the_install(self): @@ -1441,8 +1260,8 @@ def install(): reapplies = self._save_in_one_transaction(install, functools.partial(self._join_the_chassis, chassis)) (module,) = installed - self.assertEqual((self._names(module), _reapplied(reapplies)), (["pxe-3"], [module.pk])) - self.assertEqual((_journal(module), _journal(self.device)), ([], [])) + self.assertEqual((self._names(module), reapplied(reapplies)), (["pxe-3"], [module.pk])) + self.assertEqual((journal(module), journal(self.device)), ([], [])) def test_an_install_without_templates_then_a_chassis_position_change_rename_its_interface(self): bare_type = self._module_type("Bare") @@ -1457,8 +1276,8 @@ def install(): reapplies = self._save_in_one_transaction(install, self._change_the_chassis_position) (module,) = installed - self.assertEqual((self._names(module), _reapplied(reapplies)), (["et-3/0"], [module.pk])) - self.assertEqual((_journal(module), _journal(self.device)), ([], [])) + self.assertEqual((self._names(module), reapplied(reapplies)), (["et-3/0"], [module.pk])) + self.assertEqual((journal(module), journal(self.device)), ([], [])) @skipUnless(supports_module_moves(), REQUIRES_MOVE_RENAMES) def test_a_replacement_under_the_key_of_a_moved_module_is_recognised_at_the_position_of_its_install(self): @@ -1471,7 +1290,7 @@ def move_and_delete(): module.delete() def install_a_replacement(): - _give_the_next_module_id(keys[0]) + give_the_next_module_id(keys[0]) installed.append( Module.objects.create( device=self.device, module_bay=self._bay(self.device), module_type=self.adjacent_type @@ -1487,8 +1306,8 @@ def install_a_replacement(): (replacement,) = installed self.assertEqual(replacement.pk, keys[0]) - self.assertEqual((self._names(replacement), _reapplied(reapplies)), (["et-5/0"], [replacement.pk])) - self.assertEqual((_journal(replacement), _journal(self.device)), ([], [])) + self.assertEqual((self._names(replacement), reapplied(reapplies)), (["et-5/0"], [replacement.pk])) + self.assertEqual((journal(replacement), journal(self.device)), ([], [])) @skipUnless(supports_vc_position_token(), REQUIRES_VC_POSITION_TOKEN) def test_two_templates_that_claim_one_name_at_different_naming_points_rename_nothing(self): @@ -1506,7 +1325,7 @@ def install(): (module,) = installed self.assertEqual(self._names(module), ["1/0", "3/0"]) - (entry,) = _journal(module) + (entry,) = journal(module) self.assertEqual((entry.comments.count("`1/0`"), entry.comments.count("`3/0`")), (1, 1)) def test_a_failed_install_reapply_is_reported_on_each_module_and_not_by_the_device(self): @@ -1521,16 +1340,16 @@ def test_a_failed_install_reapply_is_reported_on_each_module_and_not_by_the_devi failure = f"injected {InterfaceTemplate._meta.db_table} read failure" with ( - connection.execute_wrapper(_reject_reads_of(InterfaceTemplate._meta.db_table)), + connection.execute_wrapper(reject_reads_of(InterfaceTemplate._meta.db_table)), self.assertLogs("netbox_interface_name_rules", "ERROR"), ): run_the_reapply(callbacks) for module in (first, second): - (entry,) = _journal(module) + (entry,) = journal(module) self.assertEqual(entry.kind, JournalEntryKindChoices.KIND_DANGER) self.assertEqual(entry.comments.count(failure), 1) - self.assertEqual((self._names(first), self._names(second), _journal(self.device)), (["0"], ["1"], [])) + self.assertEqual((self._names(first), self._names(second), journal(self.device)), (["0"], ["1"], [])) @skipUnless(supports_vc_position_token(), REQUIRES_VC_POSITION_TOKEN) def test_an_install_in_a_card_then_a_chassis_position_change_recognise_the_nested_raw_name(self): @@ -1542,8 +1361,8 @@ def _assert_an_install_is_named_once(self, chassis_first): module, other, reapplies = self._install_with_the_chassis_change(chassis_first) self.assertEqual((self._names(module), self._names(other)), (["et-3/0/0"], ["et-3/10/10"])) - self.assertEqual(_reapplied(reapplies), sorted((module.pk, other.pk))) - self.assertEqual((_journal(module), _journal(self.device)), ([], [])) + self.assertEqual(reapplied(reapplies), sorted((module.pk, other.pk))) + self.assertEqual((journal(module), journal(self.device)), ([], [])) def test_an_install_then_a_chassis_position_change_name_the_module_once(self): self._assert_an_install_is_named_once(chassis_first=False) @@ -1556,10 +1375,10 @@ def _assert_a_collision_is_reported_once(self, chassis_first): module, _other, reapplies = self._install_with_the_chassis_change(chassis_first) - self.assertEqual((self._names(module), _reapplied(reapplies).count(module.pk)), (["0"], 1)) - (entry,) = _journal(module) + self.assertEqual((self._names(module), reapplied(reapplies).count(module.pk)), (["0"], 1)) + (entry,) = journal(module) self.assertEqual(entry.comments.count(f"`0` to `et-3/0/0`: {TAKEN}"), 1) - self.assertEqual(_journal(self.device), []) + self.assertEqual(journal(self.device), []) def test_an_install_then_a_chassis_position_change_report_a_collision_once(self): self._assert_a_collision_is_reported_once(chassis_first=False) @@ -1650,7 +1469,7 @@ def install(): self._save_in_one_transaction(install, *device_changes) (module,) = installed - outcome = (self._names(module), [entry.comments for entry in _journal(module)]) + outcome = (self._names(module), [entry.comments for entry in journal(module)]) module.delete() return outcome @@ -1664,11 +1483,11 @@ def test_an_install_with_a_position_change_undone_ends_as_the_install_alone(self alone = self._install_outcome(module_type, bay, hand_added) self.assertEqual(alone, (names, [])) self.assertEqual(self._install_outcome(module_type, bay, hand_added, *undone), alone) - self.assertEqual(_journal(self.device), []) + self.assertEqual(journal(self.device), []) @skipUnless(supports_channelization(), REQUIRES_CHANNELIZATION) def test_a_channelized_install_with_a_position_change_undone_ends_as_the_install_alone(self): - module_type = _channelized_module_type( + module_type = channelized_module_type( self.manufacturer, f"{self.prefix} Channelized", channels=2, child_channel_ids=(1, 2) ) InterfaceNameRule.objects.create( @@ -1740,7 +1559,7 @@ def test_a_move_renames_the_moved_module_from_its_old_raw_names_and_reports_the_ callback() self.assertEqual((self._names(card), self._names(optic)), (["ge-1/2"], ["et-1/0/1"])) - (entry,) = _journal(card) + (entry,) = journal(card) self.assertIn(f"`et-1/0/1`: {STALE_BAY}", entry.comments) self.assertNotIn("ge-1/2", entry.comments) @@ -1763,7 +1582,7 @@ def test_a_nested_module_whose_bay_keeps_the_old_parent_is_not_renamed_by_anothe self._move(card, self._bay(self.device, "Bay 2")) self.assertEqual(self._names(optic), ["a-1"]) - (entry,) = _journal(card) + (entry,) = journal(card) self.assertEqual(entry.kind, JournalEntryKindChoices.KIND_WARNING) self.assertIn(f"`a-1`: {STALE_BAY}", entry.comments) @@ -1777,7 +1596,7 @@ def test_a_position_change_after_a_move_does_not_rename_a_nested_module_whose_ba self.device.save() self.assertEqual(self._names(optic), ["a-1"]) - (entry,) = _journal(self.device) + (entry,) = journal(self.device) self.assertIn(f"`a-1`: {STALE_BAY}", entry.comments) def test_a_move_to_another_device_renames_nothing_while_the_interfaces_stay_on_the_old_device(self): @@ -1788,7 +1607,7 @@ def test_a_move_to_another_device_renames_nothing_while_the_interfaces_stay_on_t interface = Interface.objects.get(module=module) self.assertEqual((interface.device, interface.name), (self.device, "et-1/0/0")) - (entry,) = _journal(module) + (entry,) = journal(module) self.assertEqual(entry.kind, JournalEntryKindChoices.KIND_WARNING) self.assertIn(f"`et-1/0/0`: {ELSEWHERE}", entry.comments) @@ -1802,11 +1621,11 @@ def test_a_position_change_of_the_new_device_renames_nothing_while_the_interface interface = Interface.objects.get(module=module) self.assertEqual((interface.device, interface.name), (self.device, "et-1/0/0")) - (entry,) = _journal(self.peer) + (entry,) = journal(self.peer) self.assertIn(f"`et-1/0/0`: {ELSEWHERE}", entry.comments) -class ModuleMoveAPITest(_MoveFixture, APITestCase): +class ModuleMoveAPITest(MoveFixture, APITestCase): """A REST API move reaches the rename trigger through NetBox's own write path.""" model = Module @@ -1832,4 +1651,4 @@ def test_patching_the_module_bay_renames_the_interfaces_for_the_new_bay(self): self.assertEqual(response.status_code, status.HTTP_200_OK, response.data) self.assertEqual(self._names(module), ["et-1/0/1"]) - self.assertEqual(_journal(module), []) + self.assertEqual(journal(module), []) diff --git a/netbox_interface_name_rules/tests/test_naming_point_sequences.py b/netbox_interface_name_rules/tests/test_naming_point_sequences.py index a9d5bec3..ee6f50d4 100644 --- a/netbox_interface_name_rules/tests/test_naming_point_sequences.py +++ b/netbox_interface_name_rules/tests/test_naming_point_sequences.py @@ -23,12 +23,8 @@ from netbox_interface_name_rules.family import supports_module_moves from netbox_interface_name_rules.models import InterfaceNameRule from netbox_interface_name_rules.naming import bay_naming_values, chassis_position -from netbox_interface_name_rules.tests.test_module_move_trigger import ( - PLAIN_TYPE, - ModuleMoveTestCase, - _module_reapplies, -) -from netbox_interface_name_rules.tests.test_vc_drift import REQUIRES_VC_POSITION_TOKEN +from netbox_interface_name_rules.tests.helpers import PLAIN_TYPE, REQUIRES_VC_POSITION_TOKEN +from netbox_interface_name_rules.tests.trigger_cases import ModuleMoveTestCase, module_reapplies # The template name after the module's code, or None for a module type without interface templates. SHAPES = { @@ -191,7 +187,7 @@ def edit_the_bays(): "bay edit": edit_the_bays, **dict.fromkeys(MOVES, move), } - with _module_reapplies() as spy: + with module_reapplies() as spy: self._save_in_one_transaction(*(saves[op] for op in sequence)) reapplies = collections.Counter(call.args[0].pk for call in spy.call_args_list) diff --git a/netbox_interface_name_rules/tests/test_prospective_families.py b/netbox_interface_name_rules/tests/test_prospective_families.py index 4d33fd87..03781e47 100644 --- a/netbox_interface_name_rules/tests/test_prospective_families.py +++ b/netbox_interface_name_rules/tests/test_prospective_families.py @@ -37,18 +37,18 @@ resolved_template_names, ) from netbox_interface_name_rules.models import InterfaceNameRule -from netbox_interface_name_rules.tests.out_of_band import rename_out_of_band -from netbox_interface_name_rules.tests.test_breakout_mode import CHANNELIZED, _plain_module_type -from netbox_interface_name_rules.tests.test_channelization import ( +from netbox_interface_name_rules.tests.helpers import ( CHANNEL_TYPE, + CHANNELIZED, PLAIN_TYPE, REQUIRES_CHANNELIZATION, + REQUIRES_NO_CHANNELIZATION, ChannelizationTestCase, - _build_device, - _channelized_module_type, + build_device, + channelized_module_type, + plain_module_type, ) - -REQUIRES_NO_CHANNELIZATION = "requires a NetBox that cannot model channelized interfaces (4.6 and older)" +from netbox_interface_name_rules.tests.out_of_band import rename_out_of_band def _projection(plan): @@ -146,8 +146,8 @@ class ProspectiveFlatPlanTest(ProspectivePlanTestCase): @classmethod def setUpTestData(cls): - manufacturer, cls.device = _build_device("ProspFlat", ["3", "4"]) - cls.module_type = _plain_module_type(manufacturer, "ProspFlat-QSFP") + manufacturer, cls.device = build_device("ProspFlat", ["3", "4"]) + cls.module_type = plain_module_type(manufacturer, "ProspFlat-QSFP") cls.rule = InterfaceNameRule.objects.create( module_type=cls.module_type, name_template="xe-0/0/{bay_position}:{channel}", @@ -225,8 +225,8 @@ class ProspectiveUnsupportedTopologyTest(ProspectivePlanTestCase): @classmethod def setUpTestData(cls): - manufacturer, cls.device = _build_device("ProspUnsup", ["3"]) - cls.module_type = _plain_module_type(manufacturer, "ProspUnsup-QSFP") + manufacturer, cls.device = build_device("ProspUnsup", ["3"]) + cls.module_type = plain_module_type(manufacturer, "ProspUnsup-QSFP") cls.rule = InterfaceNameRule.objects.create( module_type=cls.module_type, name_template="xe-0/0/{bay_position}:{channel}", @@ -258,8 +258,8 @@ class ProspectivePlansAreNotExecutableTest(ProspectivePlanTestCase): @classmethod def setUpTestData(cls): - manufacturer, cls.device = _build_device("ProspExec", ["3"]) - cls.module_type = _plain_module_type(manufacturer, "ProspExec-QSFP") + manufacturer, cls.device = build_device("ProspExec", ["3"]) + cls.module_type = plain_module_type(manufacturer, "ProspExec-QSFP") cls.rule = InterfaceNameRule.objects.create( module_type=cls.module_type, name_template="xe-0/0/{bay_position}:{channel}", @@ -292,11 +292,11 @@ class ProspectiveChannelizedPlanTest(ProspectivePlanTestCase): @classmethod def setUpTestData(cls): - manufacturer, cls.device = _build_device("ProspChan", ["3", "5", "8", "9"]) - cls.breakout_type = _channelized_module_type(manufacturer, "ProspChan-BRK") - cls.incomplete_type = _channelized_module_type(manufacturer, "ProspChan-INC", child_channel_ids=(1, 2, 3)) - cls.mismatch_type = _channelized_module_type(manufacturer, "ProspChan-MM", channels=8) - cls.simple_type = _channelized_module_type( + manufacturer, cls.device = build_device("ProspChan", ["3", "5", "8", "9"]) + cls.breakout_type = channelized_module_type(manufacturer, "ProspChan-BRK") + cls.incomplete_type = channelized_module_type(manufacturer, "ProspChan-INC", child_channel_ids=(1, 2, 3)) + cls.mismatch_type = channelized_module_type(manufacturer, "ProspChan-MM", channels=8) + cls.simple_type = channelized_module_type( manufacturer, "ProspChan-SMP", child_names={1: "{module}:1", 2: "{module}:2", 3: "{module}:3", 4: "mgmt-chan"}, @@ -370,9 +370,9 @@ class ProspectiveStructuralPlanTest(ProspectivePlanTestCase): @classmethod def setUpTestData(cls): - manufacturer, cls.device = _build_device("ProspStruct", ["3", "4"]) - cls.module_type = _plain_module_type(manufacturer, "ProspStruct-QSFP") - cls.collision_type = _plain_module_type(manufacturer, "ProspStruct-COL") + manufacturer, cls.device = build_device("ProspStruct", ["3", "4"]) + cls.module_type = plain_module_type(manufacturer, "ProspStruct-QSFP") + cls.collision_type = plain_module_type(manufacturer, "ProspStruct-COL") InterfaceTemplate.objects.create(module_type=cls.collision_type, name="et-0/0/{module}", type=PLAIN_TYPE) cls.rule = InterfaceNameRule.objects.create( module_type=cls.module_type, @@ -424,9 +424,9 @@ class ProspectiveMatchesInstalledPlanningTest(ProspectivePlanTestCase): @classmethod def setUpTestData(cls): - manufacturer, cls.device = _build_device("ProspSame", ["3", "5"]) - cls.breakout_type = _channelized_module_type(manufacturer, "ProspSame-BRK") - cls.mismatch_type = _channelized_module_type(manufacturer, "ProspSame-MM", channels=8) + manufacturer, cls.device = build_device("ProspSame", ["3", "5"]) + cls.breakout_type = channelized_module_type(manufacturer, "ProspSame-BRK") + cls.mismatch_type = channelized_module_type(manufacturer, "ProspSame-MM", channels=8) for module_type in (cls.breakout_type, cls.mismatch_type): InterfaceNameRule.objects.create( module_type=module_type, @@ -474,8 +474,8 @@ class ProspectivePreviewIsNotAppliedTest(ChannelizationTestCase): @classmethod def setUpTestData(cls): - manufacturer, cls.device = _build_device("ProspReplan", ["3"]) - cls.module_type = _plain_module_type(manufacturer, "ProspReplan-QSFP") + manufacturer, cls.device = build_device("ProspReplan", ["3"]) + cls.module_type = plain_module_type(manufacturer, "ProspReplan-QSFP") cls.rule = InterfaceNameRule.objects.create( module_type=cls.module_type, name_template="et-0/0/{bay_position}", @@ -508,8 +508,8 @@ class ProspectivePlanningIsReadOnlyTest(ChannelizationTestCase): @classmethod def setUpTestData(cls): - manufacturer, cls.device = _build_device("ProspRead", ["3"]) - cls.module_type = _plain_module_type(manufacturer, "ProspRead-QSFP") + manufacturer, cls.device = build_device("ProspRead", ["3"]) + cls.module_type = plain_module_type(manufacturer, "ProspRead-QSFP") cls.rule = InterfaceNameRule.objects.create( module_type=cls.module_type, name_template="xe-0/0/{bay_position}:{channel}", @@ -539,8 +539,8 @@ class PreviewComesFromTheFamilyPlanTest(ChannelizationTestCase): @classmethod def setUpTestData(cls): - manufacturer, cls.device = _build_device("ProspPrev", ["3"]) - cls.module_type = _plain_module_type(manufacturer, "ProspPrev-QSFP") + manufacturer, cls.device = build_device("ProspPrev", ["3"]) + cls.module_type = plain_module_type(manufacturer, "ProspPrev-QSFP") cls.rule = InterfaceNameRule.objects.create( module_type=cls.module_type, name_template="xe-0/0/{bay_position}:{channel}", @@ -581,8 +581,8 @@ class PreviewFollowsTheApplyClassificationTest(ChannelizationTestCase): @classmethod def setUpTestData(cls): - manufacturer, cls.device = _build_device("ProspClass", ["3"]) - cls.module_type = _plain_module_type(manufacturer, "ProspClass-QSFP") + manufacturer, cls.device = build_device("ProspClass", ["3"]) + cls.module_type = plain_module_type(manufacturer, "ProspClass-QSFP") cls.rule = InterfaceNameRule.objects.create( module_type=cls.module_type, name_template="et-0/0/{bay_position}", @@ -621,8 +621,8 @@ class PreviewReadsNoTemplatesForDerivableSuffixesTest(ChannelizationTestCase): @classmethod def setUpTestData(cls): - manufacturer, cls.device = _build_device("ProspQuery", ["3"]) - cls.module_type = _channelized_module_type(manufacturer, "ProspQuery-QSFP") + manufacturer, cls.device = build_device("ProspQuery", ["3"]) + cls.module_type = channelized_module_type(manufacturer, "ProspQuery-QSFP") cls.rule = InterfaceNameRule.objects.create( module_type=cls.module_type, name_template="et-0/0/{bay_position}", diff --git a/netbox_interface_name_rules/tests/test_raw_base.py b/netbox_interface_name_rules/tests/test_raw_base.py index dcb2737a..861baee1 100644 --- a/netbox_interface_name_rules/tests/test_raw_base.py +++ b/netbox_interface_name_rules/tests/test_raw_base.py @@ -23,20 +23,20 @@ supports_vc_position_token, ) from netbox_interface_name_rules.models import InterfaceNameRule -from netbox_interface_name_rules.tests.out_of_band import rename_out_of_band -from netbox_interface_name_rules.tests.test_breakout_mode import CHANNELIZED, FLAT, _plain_module_type -from netbox_interface_name_rules.tests.test_channelization import ( +from netbox_interface_name_rules.tests.helpers import ( + CHANNELIZED, + FLAT, PLAIN_TYPE, PLUGIN_LOGGER, REQUIRES_CHANNELIZATION, - _build_device, - _channelized_module_type, -) -from netbox_interface_name_rules.tests.test_vc_drift import ( REQUIRES_VC_POSITION_TOKEN, VcDriftTestCase, - _token_module_type, + build_device, + channelized_module_type, + plain_module_type, + token_module_type, ) +from netbox_interface_name_rules.tests.out_of_band import rename_out_of_band class RawBasePlainRenameTest(VcDriftTestCase): @@ -44,17 +44,17 @@ class RawBasePlainRenameTest(VcDriftTestCase): @classmethod def setUpTestData(cls): - manufacturer, cls.device = _build_device( + manufacturer, cls.device = build_device( "RawBase", ["3", "4", "5", "6", "7"], virtual_chassis=VirtualChassis.objects.create(name="rawbase-vc"), vc_position=1, ) - cls.module_type = _plain_module_type(manufacturer, "RawBase-SFP", PLAIN_TYPE) + cls.module_type = plain_module_type(manufacturer, "RawBase-SFP", PLAIN_TYPE) cls.rule = InterfaceNameRule.objects.create(module_type=cls.module_type, name_template="{base}-x") - cls.vc_type = _plain_module_type(manufacturer, "RawBase-VC", PLAIN_TYPE) + cls.vc_type = plain_module_type(manufacturer, "RawBase-VC", PLAIN_TYPE) InterfaceNameRule.objects.create(module_type=cls.vc_type, name_template="{base}.{vc_position}") - cls.fixed_type = _plain_module_type(manufacturer, "RawBase-FIXED", PLAIN_TYPE) + cls.fixed_type = plain_module_type(manufacturer, "RawBase-FIXED", PLAIN_TYPE) cls.fixed_rule = InterfaceNameRule.objects.create( module_type=cls.fixed_type, name_template="et-0/0/{bay_position}" ) @@ -66,7 +66,7 @@ def setUpTestData(cls): InterfaceTemplate.objects.create(module_type=cls.overlap_type, name="port{module}", type=PLAIN_TYPE) InterfaceTemplate.objects.create(module_type=cls.overlap_type, name="port{module}.2", type=PLAIN_TYPE) InterfaceNameRule.objects.create(module_type=cls.overlap_type, name_template="{base}.{vc_position}") - cls.flat_type = _plain_module_type(manufacturer, "RawBase-FLAT", PLAIN_TYPE) + cls.flat_type = plain_module_type(manufacturer, "RawBase-FLAT", PLAIN_TYPE) cls.flat_rule = InterfaceNameRule.objects.create( module_type=cls.flat_type, name_template="{base}:{channel}", @@ -74,17 +74,17 @@ def setUpTestData(cls): channel_count=2, channel_start=0, ) - cls.arithmetic_type = _plain_module_type(manufacturer, "RawBase-ARITH", PLAIN_TYPE) + cls.arithmetic_type = plain_module_type(manufacturer, "RawBase-ARITH", PLAIN_TYPE) InterfaceNameRule.objects.create(module_type=cls.arithmetic_type, name_template="{{base} + 100}") - cls.vc_arithmetic_type = _plain_module_type(manufacturer, "RawBase-VCARITH", PLAIN_TYPE) + cls.vc_arithmetic_type = plain_module_type(manufacturer, "RawBase-VCARITH", PLAIN_TYPE) InterfaceNameRule.objects.create( module_type=cls.vc_arithmetic_type, name_template="{{base} + 100}.{vc_position}" ) - cls.literal_marker_type = _plain_module_type(manufacturer, "RawBase-LITERAL", PLAIN_TYPE) + cls.literal_marker_type = plain_module_type(manufacturer, "RawBase-LITERAL", PLAIN_TYPE) InterfaceNameRule.objects.create( module_type=cls.literal_marker_type, name_template="{base}-InrRawBaseMark0.{vc_position}" ) - cls.assembled_marker_type = _plain_module_type(manufacturer, "RawBase-ASSEMBLED", PLAIN_TYPE) + cls.assembled_marker_type = plain_module_type(manufacturer, "RawBase-ASSEMBLED", PLAIN_TYPE) InterfaceNameRule.objects.create( module_type=cls.assembled_marker_type, name_template="{base}-InrRawBaseMark{0}.{vc_position}" ) @@ -255,8 +255,8 @@ class RawBaseChannelizedFamilyTest(VcDriftTestCase): @classmethod def setUpTestData(cls): - manufacturer, cls.device = _build_device("RawBaseChan", ["3", "7", "8"]) - cls.parent_type = _plain_module_type(manufacturer, "RawBaseChan-PARENT") + manufacturer, cls.device = build_device("RawBaseChan", ["3", "7", "8"]) + cls.parent_type = plain_module_type(manufacturer, "RawBaseChan-PARENT") cls.parent_rule = InterfaceNameRule.objects.create( module_type=cls.parent_type, name_template="xe-0/0/{bay_position}:{channel}", @@ -265,7 +265,7 @@ def setUpTestData(cls): channel_count=4, channel_start=0, ) - cls.channel_type = _plain_module_type(manufacturer, "RawBaseChan-CHANNEL") + cls.channel_type = plain_module_type(manufacturer, "RawBaseChan-CHANNEL") cls.channel_rule = InterfaceNameRule.objects.create( module_type=cls.channel_type, name_template="{base}:{channel}", @@ -274,7 +274,7 @@ def setUpTestData(cls): channel_count=4, channel_start=1, ) - cls.lockstep_type = _channelized_module_type(manufacturer, "RawBaseChan-LOCKSTEP") + cls.lockstep_type = channelized_module_type(manufacturer, "RawBaseChan-LOCKSTEP") cls.lockstep_rule = InterfaceNameRule.objects.create(module_type=cls.lockstep_type, name_template="{base}-l") def _assert_reapplies_rename_nothing(self, rule, module, bay, names): @@ -309,7 +309,7 @@ def test_a_family_whose_parent_no_template_claims_keeps_its_names(self): self.assertIn("'custom'", "\n".join(logs.output)) def test_no_family_is_built_on_an_unclaimed_plain_interface(self): - plain_type = _plain_module_type(ModuleType.objects.get(pk=self.parent_type.pk).manufacturer, "RawBaseChan-BARE") + plain_type = plain_module_type(ModuleType.objects.get(pk=self.parent_type.pk).manufacturer, "RawBaseChan-BARE") module, _ = self._install_on(self.device, plain_type, "3") rename_out_of_band(Interface.objects.get(module=module), "custom") rule = InterfaceNameRule.objects.create( @@ -339,10 +339,10 @@ class RawBaseBaseFreeParentTest(VcDriftTestCase): @classmethod def setUpTestData(cls): - manufacturer, cls.device = _build_device( + manufacturer, cls.device = build_device( "RawBaseFree", ["3"], virtual_chassis=VirtualChassis.objects.create(name="rawbasefree-vc"), vc_position=1 ) - cls.module_type = _channelized_module_type(manufacturer, "RawBaseFree-QSFP") + cls.module_type = channelized_module_type(manufacturer, "RawBaseFree-QSFP") InterfaceTemplate.objects.create(module_type=cls.module_type, name="mgmt{module}", type=PLAIN_TYPE) cls.rule = InterfaceNameRule.objects.create( module_type=cls.module_type, @@ -374,7 +374,7 @@ def test_a_parent_whose_channels_carry_two_bases_keeps_its_names(self): self.assertIn("'et-1/0/3'", "\n".join(logs.output)) def test_a_renumber_moves_a_built_parent_beside_another_template(self): - plain_type = _plain_module_type(ModuleType.objects.get(pk=self.module_type.pk).manufacturer, "RawBaseFree-SFP") + plain_type = plain_module_type(ModuleType.objects.get(pk=self.module_type.pk).manufacturer, "RawBaseFree-SFP") InterfaceTemplate.objects.create(module_type=plain_type, name="mgmt{module}", type=PLAIN_TYPE) InterfaceNameRule.objects.create( module_type=plain_type, @@ -400,13 +400,13 @@ class RawBaseFlatFamilyPreviewTest(VcDriftTestCase): @classmethod def setUpTestData(cls): - manufacturer, cls.device = _build_device( + manufacturer, cls.device = build_device( "RawBaseFlat", ["3", "4"], virtual_chassis=VirtualChassis.objects.create(name="rawbaseflat-vc"), vc_position=1, ) - cls.module_type = _token_module_type(manufacturer, "RawBaseFlat-QSFP", "xe-{vc_position:0}/0/{module}") + cls.module_type = token_module_type(manufacturer, "RawBaseFlat-QSFP", "xe-{vc_position:0}/0/{module}") cls.rule = InterfaceNameRule.objects.create( module_type=cls.module_type, name_template="brk-{base}:{channel}", @@ -447,11 +447,11 @@ class RawBaseDriftedCreationTest(VcDriftTestCase): @classmethod def setUpTestData(cls): - manufacturer, cls.device = _build_device( + manufacturer, cls.device = build_device( "RawBaseDrift", ["3"], virtual_chassis=VirtualChassis.objects.create(name="rawbasedrift-vc"), vc_position=1 ) - cls.module_type = _token_module_type(manufacturer, "RawBaseDrift-QSFP", "xe-{vc_position:0}/0/{module}") - cls.dotted_type = _channelized_module_type( + cls.module_type = token_module_type(manufacturer, "RawBaseDrift-QSFP", "xe-{vc_position:0}/0/{module}") + cls.dotted_type = channelized_module_type( manufacturer, "RawBaseDrift-DOTTED", channels=2, @@ -485,7 +485,7 @@ def test_a_blank_parent_template_keeps_the_drifted_name(self): self.assertEqual(self._names(module), ["xe-1/0/3", "xe-2/0/3:0", "xe-2/0/3:1"]) def test_a_drifted_flat_family_whose_rule_spells_a_text_sentinel_is_offered_for_conversion(self): - drift_type = _token_module_type( + drift_type = token_module_type( ModuleType.objects.get(pk=self.module_type.pk).manufacturer, "RawBaseFlat-SENT", "xe-{vc_position:0}/0/{module}", @@ -517,8 +517,8 @@ class RawBaseUnusedParentTemplateTest(VcDriftTestCase): @classmethod def setUpTestData(cls): - manufacturer, cls.device = _build_device("RawBaseUnused", ["3"]) - cls.module_type = _plain_module_type(manufacturer, "RawBaseUnused-SFP", PLAIN_TYPE) + manufacturer, cls.device = build_device("RawBaseUnused", ["3"]) + cls.module_type = plain_module_type(manufacturer, "RawBaseUnused-SFP", PLAIN_TYPE) def _rule(self, parent_name_template): """Return an unsaved flat rule, as the Build Rule tester previews one.""" diff --git a/netbox_interface_name_rules/tests/test_rename_triggers.py b/netbox_interface_name_rules/tests/test_rename_triggers.py index b2dd417f..329265cb 100644 --- a/netbox_interface_name_rules/tests/test_rename_triggers.py +++ b/netbox_interface_name_rules/tests/test_rename_triggers.py @@ -8,16 +8,20 @@ """ import gc +import pathlib import re +import tempfile from contextlib import contextmanager +from functools import partial from unittest import skipUnless from unittest.mock import patch from dcim.models import Device, Interface, InterfaceTemplate, Module, ModuleBay, VirtualChassis -from django.contrib.contenttypes.models import ContentType -from django.db import DatabaseError, DataError, IntegrityError, connection, transaction +from django.core import serializers +from django.core.management import call_command +from django.db import DEFAULT_DB_ALIAS, DatabaseError, DataError, IntegrityError, connection, transaction from django.db.models.signals import post_save -from django.test import TestCase, TransactionTestCase +from django.test import TestCase, TransactionTestCase, override_settings from django.urls import reverse from extras.choices import JournalEntryKindChoices from extras.models import JournalEntry @@ -32,21 +36,32 @@ from netbox_interface_name_rules.rename_outcomes import OutcomeKind from netbox_interface_name_rules.tests.committed_callbacks import run_the_reapply from netbox_interface_name_rules.tests.helpers import ( + PARENT_TYPE, + PLAIN_TYPE, + PLUGIN_LOGGER, + REQUIRES_CHANNELIZATION, + WriteInterfacesTo, + channelized_module_type, + interface_signal, + lock_timeout, make_device, make_device_type, make_manufacturer, make_module_bay_templates, make_module_type, + row_lock_in_another_session, + set_lock_timeout, ) from netbox_interface_name_rules.tests.out_of_band import rename_out_of_band -from netbox_interface_name_rules.tests.test_channelization import ( - PARENT_TYPE, - REQUIRES_CHANNELIZATION, - _channelized_module_type, +from netbox_interface_name_rules.tests.trigger_cases import ( + give_the_next_module_id, + journal, + module_reapplies, + previous_state_read_fails, + reject_interface_updates, + reject_reads_of, ) -PLUGIN_LOGGER = "netbox_interface_name_rules" -PLAIN_TYPE = "10gbase-x-sfpp" VIRTUAL_TYPE = "virtual" MODULE_STATE_READ = re.compile( r'SELECT "dcim_module"\."module_type_id"(?: AS "module_type_id")?, ' @@ -66,33 +81,10 @@ def _counting(entry_point): yield spy -def _module_reapplies(): - return _counting("module_rule_outcomes") - - def _device_reapplies(): return _counting("device_module_rule_outcomes") -def _journal(instance): - """Return the journal entries on *instance*, oldest first.""" - return list( - JournalEntry.objects.filter( - assigned_object_type=ContentType.objects.get_for_model(instance), assigned_object_id=instance.pk - ).order_by("pk") - ) - - -def _give_the_next_module_id(pk): - """Make the database assign *pk* to the next module that NetBox creates.""" - # NetBox before 4.7 creates no components for a module saved with an explicit pk. - with connection.cursor() as cursor: - cursor.execute( - "SELECT setval(pg_get_serial_sequence(%s, %s), %s, false)", - [Module._meta.db_table, Module._meta.pk.column, pk], - ) - - class _RenameTriggerFixture: """One device at virtual-chassis position 1 and three module types, each with its own rule. @@ -206,7 +198,7 @@ def test_a_save_that_changes_no_compared_value_causes_no_reapply(self): module = self._install() with ( - _module_reapplies() as module_reapplies, + module_reapplies() as module_calls, _device_reapplies() as device_reapplies, self.captureOnCommitCallbacks(execute=True), ): @@ -218,10 +210,10 @@ def test_a_save_that_changes_no_compared_value_causes_no_reapply(self): f"{self.prefix} New", self.device.device_type, virtual_chassis=self.virtual_chassis, vc_position=4 ) - self.assertEqual((module_reapplies.call_count, device_reapplies.call_count), (0, 0)) + self.assertEqual((module_calls.call_count, device_reapplies.call_count), (0, 0)) def test_a_module_deleted_before_commit_is_skipped(self): - with _module_reapplies() as reapplies, self.captureOnCommitCallbacks(execute=True): + with module_reapplies() as reapplies, self.captureOnCommitCallbacks(execute=True): module = Module.objects.create(device=self.device, module_bay=self._bay(), module_type=self.type_a) module.delete() @@ -307,25 +299,10 @@ def test_the_journal_entry_names_the_request_user_as_its_author(self): self._patch(self.module, {"module_type": self.type_b.pk}) - (entry,) = _journal(self.module) + (entry,) = journal(self.module) self.assertEqual(entry.created_by, self.user) -@contextmanager -def _previous_state_read_fails(statement): - """Replace the previous-state read that *statement* matches by SQL that PostgreSQL rejects.""" - replaced = [] - - def divide_by_zero(execute, sql, params, many, context): - if statement.match(sql): - replaced.append(sql) - return execute("SELECT 1/0", None, many, context) - return execute(sql, params, many, context) - - with connection.execute_wrapper(divide_by_zero): - yield replaced - - class PreviousStateReadFailureTest(RenameTriggerTestCase): """A save whose previous state cannot be read fails with the read error, on both paths.""" @@ -333,7 +310,7 @@ def test_a_module_save_fails_with_the_read_error(self): module = self._install() with ( - _previous_state_read_fails(MODULE_STATE_READ) as replaced, + previous_state_read_fails(MODULE_STATE_READ) as replaced, self.assertRaisesMessage(DataError, "division by zero"), transaction.atomic(), ): @@ -344,7 +321,7 @@ def test_a_module_save_fails_with_the_read_error(self): def test_a_device_save_fails_with_the_read_error(self): with ( - _previous_state_read_fails(DEVICE_STATE_READ) as replaced, + previous_state_read_fails(DEVICE_STATE_READ) as replaced, self.assertRaisesMessage(DataError, "division by zero"), transaction.atomic(), ): @@ -360,7 +337,7 @@ class CoalescedTriggerTest(RenameTriggerTestCase): def test_two_module_type_changes_reapply_once_with_the_final_rule(self): module = self._install() - with _module_reapplies() as reapplies, self.captureOnCommitCallbacks(execute=True), transaction.atomic(): + with module_reapplies() as reapplies, self.captureOnCommitCallbacks(execute=True), transaction.atomic(): self._change_type(module, self.type_b) self._change_type(module, self.type_c) @@ -368,7 +345,7 @@ def test_two_module_type_changes_reapply_once_with_the_final_rule(self): self.assertEqual(self._names(module), ["ge-1/0/0"]) def test_an_install_and_a_type_change_reapply_once_with_force(self): - with _module_reapplies() as reapplies, self.captureOnCommitCallbacks(execute=True), transaction.atomic(): + with module_reapplies() as reapplies, self.captureOnCommitCallbacks(execute=True), transaction.atomic(): module = Module.objects.create(device=self.device, module_bay=self._bay(), module_type=self.type_a) self._change_type(module, self.type_b) @@ -379,7 +356,7 @@ def test_a_module_type_changed_and_changed_back_is_compared_with_the_earliest_st module = self._install() rename_out_of_band(Interface.objects.get(module=module), "operator-name") - with _module_reapplies() as reapplies, self.captureOnCommitCallbacks(execute=True), transaction.atomic(): + with module_reapplies() as reapplies, self.captureOnCommitCallbacks(execute=True), transaction.atomic(): self._change_type(module, self.type_b) self._change_type(module, self.type_a) @@ -393,7 +370,7 @@ def test_an_install_under_a_reused_key_is_not_absorbed_by_a_pending_type_change( self._change_type(module, self.type_b) pk = module.pk module.delete() - _give_the_next_module_id(pk) + give_the_next_module_id(pk) module = Module.objects.create(device=self.device, module_bay=self._bay(), module_type=self.type_a) self.assertEqual(module.pk, pk) @@ -460,12 +437,6 @@ def test_a_rolled_back_savepoint_keeps_the_earlier_reapply(self): self.assertEqual(self._names(module), ["et-2/0/0"]) -def _reject_interface_updates(execute, sql, params, many, context): - if sql.lstrip().startswith('UPDATE "dcim_interface"'): - raise IntegrityError("injected reapply failure") - return execute(sql, params, many, context) - - def _reject_device_updates(execute, sql, params, many, context): if sql.lstrip().startswith('UPDATE "dcim_device"'): raise IntegrityError("injected device save failure") @@ -503,12 +474,35 @@ def test_a_post_save_sent_without_a_model_save_is_not_a_trigger(self): moved = Module.objects.get(pk=module.pk) moved.module_type = self.type_b - with _module_reapplies() as reapplies, self.captureOnCommitCallbacks(execute=True): + with module_reapplies() as reapplies, self.captureOnCommitCallbacks(execute=True): post_save.send(sender=Module, instance=moved, created=False, raw=False, using="default", update_fields=None) self.assertEqual(reapplies.call_count, 0) +class WriteAliasTest(RenameTriggerTestCase): + """A save through another alias than the write alias is unexpected state, so the trigger raises.""" + + def test_a_save_through_another_alias_raises_before_the_row_is_written(self): + self.device.vc_position = 2 + + with override_settings(DATABASE_ROUTERS=[WriteInterfacesTo("schema_elsewhere")]): + with self.assertRaisesMessage(RuntimeError, "to 'default', but the write alias is 'schema_elsewhere'"): + self.device.save() + + self.assertEqual(Device.objects.get(pk=self.device.pk).vc_position, 1) + + def test_a_post_save_sent_through_another_alias_raises(self): + """NetBox sends some post_save signals by hand, without the pre_save of a model save.""" + module = self._install() + + with override_settings(DATABASE_ROUTERS=[WriteInterfacesTo("schema_elsewhere")]): + with self.assertRaisesMessage(RuntimeError, "to 'default', but the write alias is 'schema_elsewhere'"): + post_save.send( + sender=Module, instance=module, created=True, raw=False, using="default", update_fields=None + ) + + class ReapplyFailureTest(RenameTriggerTestCase): """A reapply that fails after commit is logged, and the committed save stands.""" @@ -517,7 +511,7 @@ def test_a_failed_module_reapply_is_logged(self): with self.captureOnCommitCallbacks() as callbacks: self._change_type(module, self.type_b) - with connection.execute_wrapper(_reject_interface_updates), self.assertLogs(PLUGIN_LOGGER, "ERROR") as logs: + with connection.execute_wrapper(reject_interface_updates), self.assertLogs(PLUGIN_LOGGER, "ERROR") as logs: for callback in callbacks: callback() @@ -534,7 +528,7 @@ def test_a_failed_device_reapply_is_logged_for_modules_and_device_interfaces(sel with self.captureOnCommitCallbacks() as callbacks: self._move_to_position(2) - with connection.execute_wrapper(_reject_interface_updates), self.assertLogs(PLUGIN_LOGGER, "ERROR") as logs: + with connection.execute_wrapper(reject_interface_updates), self.assertLogs(PLUGIN_LOGGER, "ERROR") as logs: for callback in callbacks: callback() @@ -549,7 +543,7 @@ def test_a_failed_module_reload_is_logged(self): self._change_type(module, self.type_b) with ( - connection.execute_wrapper(_reject_reads_of("dcim_module")), + connection.execute_wrapper(reject_reads_of("dcim_module")), self.assertLogs(PLUGIN_LOGGER, "ERROR") as logs, ): run_the_reapply(callbacks) @@ -564,7 +558,7 @@ def test_a_module_type_read_failure_does_not_escape_the_reapply(self): self._change_type(module, self.type_b) with ( - connection.execute_wrapper(_reject_reads_of("dcim_moduletype")), + connection.execute_wrapper(reject_reads_of("dcim_moduletype")), self.assertLogs(PLUGIN_LOGGER, "ERROR") as logs, ): run_the_reapply(callbacks) @@ -580,7 +574,7 @@ def test_a_failed_device_reload_is_logged(self): self._move_to_position(2) with ( - connection.execute_wrapper(_reject_reads_of("dcim_device")), + connection.execute_wrapper(reject_reads_of("dcim_device")), self.assertLogs(PLUGIN_LOGGER, "ERROR") as logs, ): run_the_reapply(callbacks) @@ -590,17 +584,6 @@ def test_a_failed_device_reload_is_logged(self): self.assertEqual(self._names(module), ["et-1/0/0"]) -def _reject_reads_of(table): - """Return an execute wrapper that fails every read of *table*.""" - - def reject(execute, sql, params, many, context): - if sql.lstrip().startswith("SELECT") and (f'FROM "{table}"' in sql or f'JOIN "{table}"' in sql): - raise DatabaseError(f"injected {table} read failure") - return execute(sql, params, many, context) - - return reject - - def _reject_journal_writes(execute, sql, params, many, context): if sql.lstrip().startswith('INSERT INTO "extras_journalentry"'): raise DatabaseError("injected journal write failure") @@ -641,12 +624,12 @@ def test_a_name_collision_writes_one_warning_entry_on_the_module(self): with self.captureOnCommitCallbacks(execute=True): self._change_type(module, self.type_b) - (entry,) = _journal(module) + (entry,) = journal(module) self.assertEqual(entry.kind, JournalEntryKindChoices.KIND_WARNING) self.assertIn("`et-1/0/0` to `xe-1/0/0`", entry.comments) self.assertIn("already in use", entry.comments) self.assertIsNone(entry.created_by) - self.assertEqual(_journal(self.device), []) + self.assertEqual(journal(self.device), []) self.assertEqual(self._names(module), ["et-1/0/0"]) def test_the_entry_names_only_the_interfaces_left_unrenamed(self): @@ -656,7 +639,7 @@ def test_the_entry_names_only_the_interfaces_left_unrenamed(self): module = self._install(module_type) self.assertEqual(self._names(module), ["a0.1", "b0"]) - (entry,) = _journal(module) + (entry,) = journal(module) self.assertIn("`b0` to `b0.1`", entry.comments) self.assertNotIn("a0", entry.comments) @@ -667,7 +650,7 @@ def test_an_interface_no_template_claims_is_reported(self): with self.captureOnCommitCallbacks(execute=True): self._change_type(module, claiming_type) - (entry,) = _journal(module) + (entry,) = journal(module) self.assertEqual(entry.kind, JournalEntryKindChoices.KIND_WARNING) self.assertIn("`et-1/0/0`", entry.comments) self.assertIn("no single interface template claims", entry.comments) @@ -698,7 +681,7 @@ def _type_change_to_breakout(self, bay_name, name_template): return module, renamed.name def _assert_reports_each_unclaimed_top_level_interface(self, module, renamed): - (entry,) = _journal(module) + (entry,) = journal(module) self.assertEqual(entry.kind, JournalEntryKindChoices.KIND_WARNING) for name in (renamed, "mgmt-extra"): with self.subTest(name=name): @@ -726,7 +709,7 @@ def test_a_template_variable_the_device_lacks_is_reported(self): module_type=self.type_a, ) - (entry,) = _journal(module) + (entry,) = journal(module) self.assertEqual(entry.kind, JournalEntryKindChoices.KIND_WARNING) self.assertIn("`0`", entry.comments) self.assertIn("{vc_position} is not available", entry.comments) @@ -750,7 +733,7 @@ def test_a_rule_that_cannot_be_evaluated_writes_a_danger_entry(self): module = self._install(module_type) - (entry,) = _journal(module) + (entry,) = journal(module) self.assertEqual(entry.kind, JournalEntryKindChoices.KIND_DANGER) self.assertIn("`0`", entry.comments) self.assertIn("Invalid arithmetic expression", entry.comments) @@ -762,11 +745,11 @@ def test_a_failed_module_reapply_writes_a_danger_entry(self): with self.captureOnCommitCallbacks() as callbacks: self._change_type(module, self.type_b) - with connection.execute_wrapper(_reject_interface_updates), self.assertLogs(PLUGIN_LOGGER, "ERROR"): + with connection.execute_wrapper(reject_interface_updates), self.assertLogs(PLUGIN_LOGGER, "ERROR"): for callback in callbacks: callback() - (entry,) = _journal(module) + (entry,) = journal(module) self.assertEqual(entry.kind, JournalEntryKindChoices.KIND_DANGER) self.assertIn("injected reapply failure", entry.comments) @@ -779,14 +762,14 @@ def test_a_failed_device_reapply_writes_one_danger_entry_on_the_device(self): with self.captureOnCommitCallbacks() as callbacks: self._move_to_position(2) - with connection.execute_wrapper(_reject_interface_updates), self.assertLogs(PLUGIN_LOGGER, "ERROR"): + with connection.execute_wrapper(reject_interface_updates), self.assertLogs(PLUGIN_LOGGER, "ERROR"): for callback in callbacks: callback() - (entry,) = _journal(self.device) + (entry,) = journal(self.device) self.assertEqual(entry.kind, JournalEntryKindChoices.KIND_DANGER) self.assertEqual(entry.comments.count("injected reapply failure"), 2) - self.assertEqual(_journal(module), []) + self.assertEqual(journal(module), []) def test_leaving_the_virtual_chassis_renames_nothing_and_reports_the_skip(self): module = self._install() @@ -802,14 +785,14 @@ def test_leaving_the_virtual_chassis_renames_nothing_and_reports_the_skip(self): self._leave() - (entry,) = _journal(self.device) + (entry,) = journal(self.device) self.assertEqual(entry.kind, JournalEntryKindChoices.KIND_WARNING) self.assertIn("`et-1/0/0`", entry.comments) self.assertIn("`mgmt0`", entry.comments) self.assertIn("{vc_position} is not available", entry.comments) self.assertNotIn("eth0", entry.comments) self.assertNotIn("operator-name", entry.comments) - self.assertEqual(_journal(module), []) + self.assertEqual(journal(module), []) self.assertEqual((self._names(module), self._names(by_hand)), (["et-1/0/0"], ["operator-name"])) self.assertEqual( sorted(Interface.objects.filter(device=self.device, module=None).values_list("name", flat=True)), @@ -823,7 +806,7 @@ def test_a_virtual_chassis_member_without_a_position_renames_nothing_and_reports with self.captureOnCommitCallbacks(execute=True): self._move_to_position(None) - (entry,) = _journal(self.device) + (entry,) = journal(self.device) self.assertIn("`et-1/0/0`", entry.comments) self.assertNotIn("operator-name", entry.comments) self.assertEqual((self._names(module), self._names(by_hand)), (["et-1/0/0"], ["operator-name"])) @@ -845,7 +828,7 @@ def test_a_breakout_rule_the_device_cannot_evaluate_reports_the_raw_interface(se module_type=flat_type, ) - (entry,) = _journal(module) + (entry,) = journal(module) self.assertIn("`0`: {vc_position} is not available", entry.comments) self.assertEqual(self._names(module), ["0"]) @@ -868,7 +851,7 @@ def test_a_breakout_rule_the_device_cannot_evaluate_does_not_report_an_interface # No template claims "0:5", so the install scope leaves it out of every plan. Interface.objects.create(device=standalone, module=module, name="0:5", type=PLAIN_TYPE) - (entry,) = _journal(module) + (entry,) = journal(module) self.assertIn("`0`: {vc_position} is not available", entry.comments) self.assertNotIn("0:5", entry.comments) self.assertEqual(self._names(module), ["0", "0:5"]) @@ -910,7 +893,7 @@ def test_leaving_the_virtual_chassis_reports_every_member_of_a_flat_family(self) self._leave() - (entry,) = _journal(self.device) + (entry,) = journal(self.device) self.assertIn("`et-1/0:0`", entry.comments) self.assertIn("`et-1/0:1`", entry.comments) self.assertEqual(self._names(module), ["et-1/0:0", "et-1/0:1"]) @@ -930,14 +913,14 @@ def test_leaving_the_virtual_chassis_does_not_report_a_parent_that_keeps_its_nam self._leave() - (entry,) = _journal(self.device) + (entry,) = journal(self.device) self.assertIn("`xe-1/0/0:0`", entry.comments) self.assertIn("`xe-1/0/0:1`", entry.comments) self.assertNotIn("`0`", entry.comments) def _install_channelized_family(self, model, **rule_fields): """Install a module whose templates form a two-channel family, under a rule with *rule_fields*.""" - module_type = _channelized_module_type(self.type_a.manufacturer, model, channels=2, child_channel_ids=(1, 2)) + module_type = channelized_module_type(self.type_a.manufacturer, model, channels=2, child_channel_ids=(1, 2)) InterfaceNameRule.objects.create(module_type=module_type, **rule_fields) return self._install(module_type) @@ -948,7 +931,7 @@ def test_leaving_the_virtual_chassis_reports_a_parent_a_simple_rule_renames(self self._leave() - (entry,) = _journal(self.device) + (entry,) = journal(self.device) for name in ("et-1/0", "et-1/0:1", "et-1/0:2"): self.assertIn(f"`{name}`", entry.comments) @@ -964,7 +947,7 @@ def test_leaving_the_virtual_chassis_does_not_report_a_parent_a_flat_breakout_ru self._leave() - (entry,) = _journal(self.device) + (entry,) = journal(self.device) self.assertIn("`xe-1/0/0:0`", entry.comments) self.assertNotIn("`0`", entry.comments) @@ -983,7 +966,7 @@ def test_a_failure_stays_reported_when_a_lower_priority_rule_renames_the_interfa self._move_to_position(2) self.assertTrue(Interface.objects.filter(device=self.device, name="oob-2").exists()) - (entry,) = _journal(self.device) + (entry,) = journal(self.device) self.assertEqual(entry.kind, JournalEntryKindChoices.KIND_DANGER) self.assertIn("`mgmt0`", entry.comments) @@ -1007,7 +990,7 @@ def test_an_interface_a_lower_priority_rule_renames_is_not_reported(self): self._move_to_position(2) self.assertTrue(Interface.objects.filter(device=self.device, name="oob-2").exists()) - (entry,) = _journal(self.device) + (entry,) = journal(self.device) self.assertIn("`eth0` to `lan-2`", entry.comments) self.assertNotIn("mgmt0", entry.comments) @@ -1019,11 +1002,11 @@ def test_a_device_reapply_that_fails_on_one_module_keeps_the_earlier_outcomes(se with self.captureOnCommitCallbacks() as callbacks: self._move_to_position(2) - with connection.execute_wrapper(_reject_interface_updates), self.assertLogs(PLUGIN_LOGGER, "ERROR"): + with connection.execute_wrapper(reject_interface_updates), self.assertLogs(PLUGIN_LOGGER, "ERROR"): for callback in callbacks: callback() - (entry,) = _journal(self.device) + (entry,) = journal(self.device) self.assertEqual(entry.kind, JournalEntryKindChoices.KIND_DANGER) self.assertIn("`et-1/0/0` to `et-2/0/0`", entry.comments) self.assertIn("injected reapply failure", entry.comments) @@ -1043,11 +1026,11 @@ def test_a_device_reapply_that_fails_on_one_device_interface_family_keeps_the_ea with self.captureOnCommitCallbacks() as callbacks: self._move_to_position(2) - with connection.execute_wrapper(_reject_interface_updates), self.assertLogs(PLUGIN_LOGGER, "ERROR"): + with connection.execute_wrapper(reject_interface_updates), self.assertLogs(PLUGIN_LOGGER, "ERROR"): for callback in callbacks: callback() - (entry,) = _journal(self.device) + (entry,) = journal(self.device) self.assertEqual(entry.kind, JournalEntryKindChoices.KIND_DANGER) self.assertIn("`mgmt0` to `mgmt-2`", entry.comments) self.assertIn("injected reapply failure", entry.comments) @@ -1062,11 +1045,11 @@ def test_a_module_reapply_that_fails_on_one_family_keeps_the_earlier_outcomes(se with self.captureOnCommitCallbacks() as callbacks: module = Module.objects.create(device=self.device, module_bay=self._bay(), module_type=module_type) - with connection.execute_wrapper(_reject_interface_updates), self.assertLogs(PLUGIN_LOGGER, "ERROR"): + with connection.execute_wrapper(reject_interface_updates), self.assertLogs(PLUGIN_LOGGER, "ERROR"): for callback in callbacks: callback() - (entry,) = _journal(module) + (entry,) = journal(module) self.assertEqual(entry.kind, JournalEntryKindChoices.KIND_DANGER) self.assertIn("`a0` to `a0.1`", entry.comments) self.assertIn("injected reapply failure", entry.comments) @@ -1103,7 +1086,7 @@ def test_a_failed_journal_write_is_logged_and_does_not_raise(self): callback() self.assertEqual([str(record.exc_info[1]) for record in logs.records], ["injected journal write failure"]) - self.assertEqual(_journal(module), []) + self.assertEqual(journal(module), []) self.assertEqual(Module.objects.get(pk=module.pk).module_type, self.type_b) @@ -1130,6 +1113,20 @@ def test_without_an_open_transaction_the_module_reapply_reads_the_saved_row(self self.assertEqual(self._names(self.module), ["xe-1/0/0"]) + def test_a_fixture_load_of_a_module_reapplies_as_an_install(self): + """Django's loaddata saves each row raw, and a raw save is a rename trigger too.""" + interface = rename_out_of_band(Interface.objects.get(module=self.module), "0") + fixture = serializers.serialize("json", [self.module, interface]) + pk = self.module.pk + self.module.delete() + + with tempfile.TemporaryDirectory() as directory: + path = pathlib.Path(directory) / "module.json" + path.write_text(fixture, encoding="utf-8") + call_command("loaddata", str(path), verbosity=0) + + self.assertEqual(self._names(Module.objects.get(pk=pk)), ["et-1/0/0"]) + def test_a_rolled_back_transaction_does_not_suppress_the_next_trigger(self): with self.assertRaises(RuntimeError), transaction.atomic(): self.device.vc_position = 2 @@ -1142,3 +1139,49 @@ def test_a_rolled_back_transaction_does_not_suppress_the_next_trigger(self): self.device.save() self.assertEqual(self._names(self.module), ["et-3/0/0"]) + + +@skipUnless(supports_channelization(), REQUIRES_CHANNELIZATION) +class ReconciliationFailureJournalTest(TransactionTestCase): + """A channel reconciliation that fails after the commit names each channel and the name it kept.""" + + def setUp(self): + manufacturer = make_manufacturer("ReconJournal") + device_type = make_device_type(manufacturer, "ReconJournal") + make_module_bay_templates(device_type, ("Bay 0", "Bay 1")) + self.device = make_device("ReconJournal", device_type) + self.module_type = channelized_module_type(manufacturer, "ReconJournal-QSFP") + InterfaceNameRule.objects.create( + module_type=self.module_type, + name_template="xe-0/0/{bay_position}:{channel}", + parent_name_template="et-0/0/{bay_position}", + breakout_mode=BreakoutModeChoices.CHANNELIZED, + channel_count=4, + channel_start=0, + ) + # Channel 2 cannot take its target, so the rule keeps the name NetBox gave it. + Interface.objects.create(device=self.device, name="xe-0/0/1:1", type=PLAIN_TYPE) + # On main the plugin sets no timeout, and without one the reconciliation would wait for the lock. + self.addCleanup(set_lock_timeout, DEFAULT_DB_ALIAS, lock_timeout(DEFAULT_DB_ALIAS)) + set_lock_timeout(DEFAULT_DB_ALIAS, "1s") + + def test_the_journal_names_each_channel_and_its_kept_name(self): + """A second session locks the kept channel after NetBox's cascade renamed it, before the reconciliation.""" + bay = ModuleBay.objects.get(device=self.device, name="Bay 1") + + with row_lock_in_another_session("default") as lock: + + def lock_after_the_cascade(sender, instance, **kwargs): + if instance.name == "et-0/0/1:2": + transaction.on_commit(partial(lock, instance.pk), using=instance._state.db) + + with interface_signal(post_save, lock_after_the_cascade), transaction.atomic(): + module = Module.objects.create(device=self.device, module_bay=bay, module_type=self.module_type) + + [entry] = journal(module) + self.assertEqual(entry.kind, JournalEntryKindChoices.KIND_DANGER) + self.assertIn("Rename each channel back: `et-0/0/1:2` to `1:2`", entry.comments) + self.assertEqual( + sorted(Interface.objects.filter(module=module).values_list("name", flat=True)), + ["et-0/0/1", "et-0/0/1:2", "xe-0/0/1:0", "xe-0/0/1:2", "xe-0/0/1:3"], + ) diff --git a/netbox_interface_name_rules/tests/test_rule_validation_agreement.py b/netbox_interface_name_rules/tests/test_rule_validation_agreement.py index 98ff31c1..2937cb5f 100644 --- a/netbox_interface_name_rules/tests/test_rule_validation_agreement.py +++ b/netbox_interface_name_rules/tests/test_rule_validation_agreement.py @@ -263,6 +263,45 @@ def test_a_targeted_save_validates_on_the_database_the_router_writes_to(self): self.assertEqual(InterfaceNameRule.objects.get(pk=rule.pk).name_template, "xe-0/{bay_position}") + def test_a_targeted_save_through_an_alias_other_than_the_write_alias_is_refused(self): + """A rule save writes through the write alias alone; on main that is ``default``.""" + rule = InterfaceNameRule.objects.create(module_type=self.module_type, name_template="xe-{bay_position}") + rule.name_template = "xe-0/{bay_position}" + + with self.assertRaisesMessage(RuntimeError, "'schema_elsewhere'"): + rule.save(using="schema_elsewhere", update_fields=["name_template"]) + + self.assertEqual(InterfaceNameRule.objects.get(pk=rule.pk).name_template, "xe-{bay_position}") + + def test_a_full_save_through_an_alias_other_than_the_write_alias_is_refused(self): + rule = InterfaceNameRule.objects.create(module_type=self.module_type, name_template="xe-{bay_position}") + rule.name_template = "xe-0/{bay_position}" + + with self.assertRaisesMessage(RuntimeError, "'schema_elsewhere'"): + rule.save(using="schema_elsewhere") + + self.assertEqual(InterfaceNameRule.objects.get(pk=rule.pk).name_template, "xe-{bay_position}") + + def test_a_save_of_an_unrelated_field_through_an_alias_other_than_the_write_alias_is_refused(self): + rule = InterfaceNameRule.objects.create(module_type=self.module_type, name_template="xe-{bay_position}") + rule.description = "after" + + with self.assertRaisesMessage(RuntimeError, "'schema_elsewhere'"): + rule.save(using="schema_elsewhere", update_fields=["description"]) + + self.assertEqual(InterfaceNameRule.objects.get(pk=rule.pk).description, "") + + def test_a_targeted_save_whose_routed_alias_is_outside_the_write_scope_is_refused(self): + """The router may send a rule where the write scope of an interface write does not reach.""" + rule = InterfaceNameRule.objects.create(module_type=self.module_type, name_template="xe-{bay_position}") + rule.name_template = "xe-0/{bay_position}" + + with override_settings(DATABASE_ROUTERS=[WriteRulesTo("schema_elsewhere")]): + with self.assertRaisesMessage(RuntimeError, "'schema_elsewhere', outside the write scope ('default',)"): + rule.save(update_fields=["name_template"]) + + self.assertEqual(InterfaceNameRule.objects.get(pk=rule.pk).name_template, "xe-{bay_position}") + def test_a_deferred_save_keeps_a_concurrent_change_to_a_field_it_did_not_load(self): """Validation must not load the deferred fields, which Django then writes back over a concurrent change.""" rule = InterfaceNameRule.objects.create( @@ -460,6 +499,16 @@ def db_for_write(self, model, **hints): return "default" +class WriteRulesTo: + """A database router that sends each write of a rule to one alias.""" + + def __init__(self, alias): + self.alias = alias + + def db_for_write(self, model, **hints): + return self.alias if model is InterfaceNameRule else None + + class RefuseImplicitMigrationDatabase: """Refuse migration queries that do not select a database explicitly.""" diff --git a/netbox_interface_name_rules/tests/test_rules.py b/netbox_interface_name_rules/tests/test_rules.py index 441c2ea5..2d157c25 100644 --- a/netbox_interface_name_rules/tests/test_rules.py +++ b/netbox_interface_name_rules/tests/test_rules.py @@ -5,6 +5,7 @@ These tests create real DB objects and exercise the full engine pipeline. """ +import dataclasses import threading from dcim.models import ( @@ -17,6 +18,7 @@ ModuleType, Platform, ) +from django.db import DEFAULT_DB_ALIAS from django.test import TestCase from netbox_interface_name_rules.engine import ( @@ -303,10 +305,10 @@ def setUp(self): # and a stale snapshot from a prior method can't be reused for a same-content rule set. from netbox_interface_name_rules import rule_selection - rule_selection._RULE_CACHE.update({"version": None, "exact": (), "regex": (), "memo": {}}) + rule_selection._RULE_CACHE = rule_selection._empty_rule_cache() rule_selection._pin.depth = 0 rule_selection._pin.primed = False - for attr in ("exact", "regex", "memo"): + for attr in ("alias", "exact", "regex", "memo"): rule_selection._pin.__dict__.pop(attr, None) def test_repeated_calls_do_not_re_query_rules(self): @@ -439,12 +441,12 @@ def test_compensating_fk_swap_changes_fingerprint(self): rule_a = InterfaceNameRule.objects.create(module_type=mt_a, device_type=dt_x, name_template="a{bay_position}") rule_b = InterfaceNameRule.objects.create(module_type=mt_b, device_type=dt_y, name_template="b{bay_position}") - before = rule_selection._enabled_rules_version() + before = rule_selection._enabled_rules_version(DEFAULT_DB_ALIAS) # Swap device_type between the two rules via bulk .update(): SUM(device_type) is unchanged # (x+y == y+x), count unchanged, last_updated unbumped. Only the per-rule pairing differs. InterfaceNameRule.objects.filter(pk=rule_a.pk).update(device_type=dt_y) InterfaceNameRule.objects.filter(pk=rule_b.pk).update(device_type=dt_x) - after = rule_selection._enabled_rules_version() + after = rule_selection._enabled_rules_version(DEFAULT_DB_ALIAS) self.assertNotEqual(before, after, "fingerprint collided on a compensating FK swap (Sum-aggregate weakness)") @@ -497,7 +499,7 @@ def test_memo_is_bounded(self): # would slip past an end-of-loop `<=` snapshot (which only sees the post-clear size) but # trips here on the very insert that crosses the cap. self.assertLessEqual( - len(rule_selection._RULE_CACHE["memo"]), + len(rule_selection._RULE_CACHE.memo), rule_selection._MEMO_MAX, f"memo exceeded the cap mid-insertion after {i} contexts", ) @@ -505,7 +507,7 @@ def test_memo_is_bounded(self): # And the eviction actually fired: the final size is below the number of distinct contexts, # so the test isn't vacuously green on a memo that simply never reached the cap. self.assertLess( - len(rule_selection._RULE_CACHE["memo"]), + len(rule_selection._RULE_CACHE.memo), len(module_types), "memo never evicted — the cap was never exercised", ) @@ -560,7 +562,7 @@ def test_pinned_block_holds_snapshot_across_concurrent_cache_reload(self): different rule set — renaming one device's modules with mixed rule versions. This exercises the guarantee across an actual thread: the worker must NOT inherit this thread's pin (``_pin`` is ``threading.local``), and its reload — published the way ``_get_enabled_rules`` does, by - rebinding ``_RULE_CACHE`` to a fresh dict — must leave our primed snapshot untouched. A plain + rebinding ``_RULE_CACHE`` to a new snapshot — must leave our primed snapshot untouched. A plain global ``_pin`` would leak the pin into the worker and pass the old same-thread test, but fail here. """ from netbox_interface_name_rules import rule_selection @@ -576,12 +578,9 @@ def concurrent_reload(): # way _get_enabled_rules() does: one atomic rebind of the module global. (No DB access — a # separate thread has its own connection and cannot see this TestCase's uncommitted rows.) worker_pin_depth.append(getattr(rule_selection._pin, "depth", 0)) - rule_selection._RULE_CACHE = { - "version": "concurrent-reload-other-version", - "exact": (), - "regex": (), - "memo": {}, - } + rule_selection._RULE_CACHE = rule_selection._RuleSnapshot( + alias=DEFAULT_DB_ALIAS, version="concurrent-reload-other-version", exact=(), regex=(), memo={} + ) try: with pinned_rule_cache(): @@ -594,7 +593,7 @@ def concurrent_reload(): # The worker did not inherit our pin — proves _pin is per-thread, not a shared global. self.assertEqual(worker_pin_depth, [0], "pin leaked across threads — _pin is not thread-local") # ...and its reload really did replace the shared cache. - self.assertEqual(rule_selection._RULE_CACHE["version"], "concurrent-reload-other-version") + self.assertEqual(rule_selection._RULE_CACHE.version, "concurrent-reload-other-version") # Still inside the pin: must serve the snapshot captured at entry, not the worker's reload. self.assertEqual( @@ -605,14 +604,17 @@ def concurrent_reload(): finally: # This test deliberately rebinds the module global from another thread; restore a clean # sentinel so the simulated reload can't leak into sibling tests even if setUp() is weakened. - rule_selection._RULE_CACHE = {"version": None, "exact": (), "regex": (), "memo": {}} + rule_selection._RULE_CACHE = rule_selection._empty_rule_cache() + + # The restored cache must serve an unpinned lookup, so a later test cannot inherit a broken cache. + self.assertEqual(find_matching_rule(self.module_type, None, self.device_type), universal) - def test_reload_publishes_a_fresh_cache_dict_atomically(self): - """A version change rebinds _RULE_CACHE to a new dict instead of mutating the old one in place. + def test_reload_publishes_a_fresh_cache_snapshot_atomically(self): + """A version change rebinds _RULE_CACHE to a new snapshot instead of mutating the old one in place. - This is what makes the unpinned three-key read a consistent snapshot: a reader that grabbed the + This is what makes the unpinned read a consistent snapshot: a reader that grabbed the cache before a concurrent reload keeps one whole rule-set version, never exact from V1 paired - with memo from V2. We assert the published dict is a *new object* and that a reference captured + with memo from V2. We assert the published snapshot is a *new object* and that a reference captured before the reload is left untouched — the in-place mutation the previous code did would fail both. """ from netbox_interface_name_rules import rule_selection @@ -621,8 +623,8 @@ def test_reload_publishes_a_fresh_cache_dict_atomically(self): find_matching_rule(self.module_type, None, self.device_type) # prime version 1 snap = rule_selection._RULE_CACHE - snap_version = snap["version"] - snap_exact = snap["exact"] + snap_version = snap.version + snap_exact = snap.exact # Change the rule set so the next lookup must reload to a new version. InterfaceNameRule.objects.create( @@ -633,10 +635,10 @@ def test_reload_publishes_a_fresh_cache_dict_atomically(self): self.assertIsNot( rule_selection._RULE_CACHE, snap, - "reload mutated the cache dict in place instead of publishing a new one", + "reload mutated the cache snapshot in place instead of publishing a new one", ) - self.assertEqual(snap["version"], snap_version, "a reload mutated a previously-published cache dict") - self.assertEqual(snap["exact"], snap_exact, "the captured exact snapshot changed under a concurrent reload") + self.assertEqual(snap.version, snap_version, "a reload mutated a previously-published cache snapshot") + self.assertEqual(snap.exact, snap_exact, "the captured exact snapshot changed under a concurrent reload") def test_find_matching_rule_survives_concurrent_memo_clear(self): """A memo cleared between the membership check and the lookup must not raise KeyError. @@ -662,7 +664,8 @@ def __contains__(self, key): return present # Same version → no reload; the racing memo is what find_matching_rule reads next. - rule_selection._RULE_CACHE["memo"] = _RacingMemo(rule_selection._RULE_CACHE["memo"]) + cache = rule_selection._RULE_CACHE + rule_selection._RULE_CACHE = dataclasses.replace(cache, memo=_RacingMemo(cache.memo)) self.assertEqual( find_matching_rule(self.module_type, None, self.device_type), @@ -683,14 +686,14 @@ def test_pinned_block_uses_a_private_memo_copy(self): self.assertIsNot( rule_selection._pin.memo, - rule_selection._RULE_CACHE["memo"], + rule_selection._RULE_CACHE.memo, "pinned memo aliases the shared cache memo instead of holding a private copy", ) # An unpinned thread clearing the shared memo at the cap must not disturb the pinned # batch. _pin is a threading.local, so the worker holds no pin of its own. The target is # a plain dict.clear(), so the worker touches no ORM and needs no connection. - clear = threading.Thread(target=rule_selection._RULE_CACHE["memo"].clear) + clear = threading.Thread(target=rule_selection._RULE_CACHE.memo.clear) clear.start() clear.join() @@ -702,7 +705,7 @@ def test_pinned_block_uses_a_private_memo_copy(self): # memo instead of its private copy would return the decoy rather than recomputing, so # this is what separates "the copy was used" from "the answer was recomputed". decoy = InterfaceNameRule(name_template="decoy-must-not-be-returned") - rule_selection._RULE_CACHE["memo"].update(dict.fromkeys(rule_selection._pin.memo, decoy)) + rule_selection._RULE_CACHE.memo.update(dict.fromkeys(rule_selection._pin.memo, decoy)) self.assertEqual( find_matching_rule(self.module_type, None, self.device_type), rule, @@ -730,7 +733,7 @@ def test_fingerprint_resists_separator_injection_in_text_fields(self): r2 = InterfaceNameRule.objects.create( applies_to_device_interfaces=True, module_type_pattern="p2", name_template="b" ) - fp_two_rules = rule_selection._enabled_rules_version() + fp_two_rules = rule_selection._enabled_rules_version(DEFAULT_DB_ALIAS) # Forge a one-rule set whose name_template embeds r1's trailing columns, a row separator, and # r2's columns up to its name_template. Column order: id, module_type_id, is_regex, pattern, @@ -740,7 +743,7 @@ def test_fingerprint_resists_separator_injection_in_text_fields(self): forged_name_template = field_sep.join(["a", "0", "0", "true"]) + row_sep + r2_cells_through_name InterfaceNameRule.objects.filter(pk=r2.pk).delete() InterfaceNameRule.objects.filter(pk=r1.pk).update(name_template=forged_name_template) - fp_forged_one_rule = rule_selection._enabled_rules_version() + fp_forged_one_rule = rule_selection._enabled_rules_version(DEFAULT_DB_ALIAS) self.assertNotEqual( fp_two_rules, diff --git a/netbox_interface_name_rules/tests/test_snapshot_guard.py b/netbox_interface_name_rules/tests/test_snapshot_guard.py index c8f94f8d..41bf3b18 100644 --- a/netbox_interface_name_rules/tests/test_snapshot_guard.py +++ b/netbox_interface_name_rules/tests/test_snapshot_guard.py @@ -6,6 +6,7 @@ plugin code, of NetBox code or of test code exactly as it sees the real ones. """ +import functools import types from dcim.choices import InterfaceModeChoices @@ -245,17 +246,24 @@ def test_a_change_that_the_row_s_own_save_makes_is_part_of_that_save(self): def test_a_tag_change_in_the_request_that_created_the_row_joins_the_create_record(self): site = Site(name="Guard Created", slug="guard-created") - run_as_job_user(make_job("GuardOne"), lambda: (plugin_save(site), plugin_add_tags(site, self.tag))) + run_as_job_user( + make_job("GuardOne"), + lambda: (plugin_save(site), plugin_add_tags(site, self.tag)), + branch_schema_id=None, + ) self.assertEqual(snapshot_guard.take_violations(), []) def test_a_tag_change_in_a_later_request_needs_a_snapshot(self): """NetBox merges an M2M change only into a record of the same request, so the later one needs a before-state.""" site = Site(name="Guard Created", slug="guard-created") - run_as_job_user(make_job("GuardFirst"), lambda: plugin_save(site)) + run_as_job_user(make_job("GuardFirst"), lambda: plugin_save(site), branch_schema_id=None) self.assert_refused( - NO_SNAPSHOT, run_as_job_user, make_job("GuardSecond"), lambda: plugin_add_tags(site, self.tag) + NO_SNAPSHOT, + functools.partial(run_as_job_user, branch_schema_id=None), + make_job("GuardSecond"), + lambda: plugin_add_tags(site, self.tag), ) def test_a_tag_change_of_a_row_that_a_test_created_needs_a_snapshot(self): diff --git a/netbox_interface_name_rules/tests/test_structural_families.py b/netbox_interface_name_rules/tests/test_structural_families.py index 731caecc..1649b06b 100644 --- a/netbox_interface_name_rules/tests/test_structural_families.py +++ b/netbox_interface_name_rules/tests/test_structural_families.py @@ -29,18 +29,19 @@ ) from netbox_interface_name_rules.family import names as family_names from netbox_interface_name_rules.models import InterfaceNameRule -from netbox_interface_name_rules.tests.helpers import make_device -from netbox_interface_name_rules.tests.out_of_band import rename_out_of_band -from netbox_interface_name_rules.tests.test_breakout_mode import CHANNELIZED, _plain_module_type -from netbox_interface_name_rules.tests.test_channelization import ( +from netbox_interface_name_rules.tests.helpers import ( + CHANNELIZED, PLAIN_TYPE, PLUGIN_LOGGER, REQUIRES_CHANNELIZATION, + REQUIRES_NO_CHANNELIZATION, ChannelizationTestCase, - _build_device, + build_device, + make_device, + plain_module_type, ) - -REQUIRES_NO_CHANNELIZATION = "requires a NetBox that cannot model channelized interfaces (4.6 and older)" +from netbox_interface_name_rules.tests.out_of_band import rename_out_of_band +from netbox_interface_name_rules.transactions import write_scope class StructuralFamilyTestCase(ChannelizationTestCase): @@ -51,8 +52,8 @@ class StructuralFamilyTestCase(ChannelizationTestCase): @classmethod def setUpTestData(cls): - manufacturer, cls.device = _build_device(cls.PREFIX, ["3", "4"]) - cls.module_type = _plain_module_type(manufacturer, f"{cls.PREFIX}-QSFP") + manufacturer, cls.device = build_device(cls.PREFIX, ["3", "4"]) + cls.module_type = plain_module_type(manufacturer, f"{cls.PREFIX}-QSFP") cls.rule = InterfaceNameRule.objects.create( module_type=cls.module_type, name_template=cls.NAME_TEMPLATE, @@ -141,7 +142,7 @@ def _cascade(child, name): def test_the_intended_name_is_restored_after_the_cascade(self): child = self._interface("et-0/0/3:1") - with self.captureOnCommitCallbacks(execute=True) as callbacks: + with self.captureOnCommitCallbacks(execute=True) as callbacks, write_scope(): family_names.reconcile_after_parent_cascade("et-0/0/3", "xe-0/0/3", ((child.pk, 1, "et-0/0/3:1"),)) self._cascade(child, "xe-0/0/3:1") @@ -152,7 +153,7 @@ def test_the_intended_name_is_restored_after_the_cascade(self): def test_a_parent_that_kept_its_name_registers_no_callback(self): child = self._interface("et-0/0/3:1") - with self.captureOnCommitCallbacks(execute=True) as callbacks: + with self.captureOnCommitCallbacks(execute=True) as callbacks, write_scope(): family_names.reconcile_after_parent_cascade("et-0/0/3", "et-0/0/3", ((child.pk, 1, "et-0/0/3:1"),)) self.assertEqual(callbacks, []) @@ -160,7 +161,7 @@ def test_a_parent_that_kept_its_name_registers_no_callback(self): def test_a_channel_the_cascade_will_not_touch_registers_no_callback(self): child = self._interface("ge-0/0/3-1") - with self.captureOnCommitCallbacks(execute=True) as callbacks: + with self.captureOnCommitCallbacks(execute=True) as callbacks, write_scope(): family_names.reconcile_after_parent_cascade("et-0/0/3", "xe-0/0/3", ((child.pk, 1, "ge-0/0/3-1"),)) self.assertEqual(callbacks, []) @@ -168,7 +169,7 @@ def test_a_channel_the_cascade_will_not_touch_registers_no_callback(self): def test_a_channel_the_cascade_left_alone_keeps_its_intended_name(self): child = self._interface("et-0/0/3:1") - with self.captureOnCommitCallbacks(execute=True): + with self.captureOnCommitCallbacks(execute=True), write_scope(): family_names.reconcile_after_parent_cascade("et-0/0/3", "xe-0/0/3", ((child.pk, 1, "et-0/0/3:1"),)) child.refresh_from_db() @@ -178,7 +179,7 @@ def test_a_channel_moved_to_an_unexpected_name_is_left_alone(self): child = self._interface("et-0/0/3:1") with self.assertLogs(PLUGIN_LOGGER, level="WARNING") as logs: - with self.captureOnCommitCallbacks(execute=True): + with self.captureOnCommitCallbacks(execute=True), write_scope(): family_names.reconcile_after_parent_cascade("et-0/0/3", "xe-0/0/3", ((child.pk, 1, "et-0/0/3:1"),)) self._cascade(child, "someone-else-renamed-it") @@ -189,7 +190,7 @@ def test_a_channel_moved_to_an_unexpected_name_is_left_alone(self): def test_a_name_taken_since_the_cascade_is_not_reclaimed(self): child = self._interface("et-0/0/3:1") - with self.assertLogs(PLUGIN_LOGGER, level="ERROR"), self.captureOnCommitCallbacks(execute=True): + with self.assertLogs(PLUGIN_LOGGER, level="ERROR"), self.captureOnCommitCallbacks(execute=True), write_scope(): family_names.reconcile_after_parent_cascade("et-0/0/3", "xe-0/0/3", ((child.pk, 1, "et-0/0/3:1"),)) self._cascade(child, "xe-0/0/3:1") occupant = self._interface("et-0/0/3:1") @@ -215,7 +216,7 @@ def test_the_reconciliation_locks_only_interfaces_in_primary_key_order(self): self.assertNotIn("dcim_device", locking[0].split("FOR UPDATE")[1]) self.assertIn('ORDER BY "dcim_interface"."id" ASC', locking[0]) - def test_an_unrelated_integrity_failure_propagates(self): + def test_an_unrelated_integrity_failure_names_the_channel_and_its_kept_name(self): interface = self._interface("cascade-name") def reject_interface_update(execute, sql, params, many, context): @@ -224,11 +225,15 @@ def reject_interface_update(execute, sql, params, many, context): return execute(sql, params, many, context) with connection.execute_wrapper(reject_interface_update): - with self.assertRaisesMessage(IntegrityError, "injected deferred database failure"): + with self.assertRaisesMessage( + family_names.ChannelReconciliationError, "`cascade-name` to `final-name`" + ) as raised: family_names.restore_deferred_channel_names( ((interface.pk, "final-name", "cascade-name"),), ) + self.assertEqual(str(raised.exception.__cause__), "injected deferred database failure") + @skipUnless(supports_channelization(), REQUIRES_CHANNELIZATION) class StructuralFamilyPlanTest(StructuralFamilyTestCase): diff --git a/netbox_interface_name_rules/tests/test_transactions.py b/netbox_interface_name_rules/tests/test_transactions.py index e403d14c..37113aff 100644 --- a/netbox_interface_name_rules/tests/test_transactions.py +++ b/netbox_interface_name_rules/tests/test_transactions.py @@ -3,8 +3,7 @@ """An atomic block keeps the NetBox events it queued only when it commits, coalesced as NetBox does.""" import contextvars -import uuid -from contextlib import contextmanager, nullcontext +from contextlib import nullcontext from itertools import product from unittest import skipUnless @@ -12,15 +11,16 @@ from dcim.models import Interface, Site from django.contrib.auth import get_user_model from django.db import IntegrityError, connection, transaction -from django.http import HttpRequest -from django.test import SimpleTestCase, TestCase, TransactionTestCase +from django.test import SimpleTestCase, TestCase, TransactionTestCase, override_settings +from django.test.utils import CaptureQueriesContext from extras import events as netbox_events from extras.events import serialize_for_event from extras.models import Tag -from netbox.context import current_request, events_queue +from netbox.context import events_queue from netbox_interface_name_rules.jobs import run_as_job_user from netbox_interface_name_rules.tests.helpers import ( + WriteInterfacesTo, empty_the_webhook_queue, make_device, make_device_type, @@ -29,8 +29,9 @@ make_manufacturer, queued_webhook_jobs, queued_webhooks, + request_context, ) -from netbox_interface_name_rules.transactions import atomic_with_events +from netbox_interface_name_rules.transactions import atomic_with_events, on_commit, write_scope def _describe(interface, description): @@ -54,7 +55,7 @@ def setUp(self): def _run_as_a_request(self, body): # django-rq enqueues a webhook when the transaction commits. with self.captureOnCommitCallbacks(execute=True): - run_as_job_user(make_job("TxEvents"), body) + run_as_job_user(make_job("TxEvents"), body, branch_schema_id=None) def test_a_change_before_and_inside_a_committed_block_is_one_event(self): def body(): @@ -167,10 +168,10 @@ def test_a_delete_before_a_block_that_recreates_the_row_and_rolls_back_stays_a_d def body(): interface.delete() - with atomic_with_events(): + with atomic_with_events() as block: interface.pk = pk interface.save() - transaction.set_rollback(True) + block.set_rollback() self._run_as_a_request(body) @@ -189,19 +190,19 @@ def body(): _describe(self.interface, "before the block") with connection.cursor() as cursor: cursor.execute('DELETE FROM "dcim_interface" WHERE id = %s', [pk]) - with atomic_with_events(): + with atomic_with_events() as block: # Django inserts the row again when the update finds none. _describe(self.interface, "inside the block") - transaction.set_rollback(True) + block.set_rollback() with self.assertRaisesMessage(RuntimeError, f"dcim.Interface {pk} has a queued event but no row"): - run_as_job_user(make_job("TxEventsGone"), body) + run_as_job_user(make_job("TxEventsGone"), body, branch_schema_id=None) def test_a_block_that_sets_rollback_drops_its_events(self): def body(): - with atomic_with_events(): + with atomic_with_events() as block: _describe(self.interface, "inside the block") - transaction.set_rollback(True) + block.set_rollback() self._run_as_a_request(body) @@ -222,12 +223,12 @@ class _BlockRollbackError(Exception): """Raised inside a block to roll it back; the level around the block catches it.""" -def _finish(outcome): - """End a block as *outcome* says: return, raise, or mark it for rollback.""" +def _finish(outcome, block): + """End *block* as *outcome* says: return, raise, or mark it for rollback.""" if outcome == "raises": raise _BlockRollbackError if outcome == "sets rollback": - transaction.set_rollback(True) + block.set_rollback() # Each step changes the row the case starts from, through the same instance or a second one, or creates a row. @@ -247,21 +248,6 @@ def generated_cases(): yield event_before, levels -@contextmanager -def request_context(user): - """Set the request and the event queue as NetBox's event_tracking does, without its flush.""" - request = HttpRequest() - request.user = user - request.id = uuid.uuid4() - request_token = current_request.set(request) - queue_token = events_queue.set({}) - try: - yield - finally: - events_queue.reset(queue_token) - current_request.reset(request_token) - - def _event_key(pk): return f"dcim.interface:{pk}" @@ -334,12 +320,12 @@ def _change(self, number, existing, ending, callbacks): _describe(existing, "before the block") new_pk = None try: - with atomic_with_events(): + with atomic_with_events() as block: _describe(existing, "inside the block") new_pk = Interface.objects.create( device=self.device, name=f"txoutcome-new-{number}", type="1000base-t" ).pk - self._end(number, ending, callbacks) + self._end(number, ending, callbacks, block) except (_BlockRollbackError, IntegrityError): return new_pk except _CallbackError: @@ -347,11 +333,11 @@ def _change(self, number, existing, ending, callbacks): return new_pk @staticmethod - def _end(number, ending, callbacks): + def _end(number, ending, callbacks, block): if ending == "raises": raise _BlockRollbackError if ending == "sets rollback": - transaction.set_rollback(True) + block.set_rollback() elif ending == "fails at COMMIT": # No tenant has this ID; PostgreSQL checks the deferred foreign key at COMMIT. Site.objects.create( @@ -413,11 +399,11 @@ def _check_case(self, event_before, levels): def _run_level(self, levels, depth, row, created): (operation, instance), outcome = levels[depth] try: - with atomic_with_events(): + with atomic_with_events() as block: self._apply(operation, row if instance == "same instance" else None, row.pk, depth, created) if depth + 1 < len(levels): self._run_level(levels, depth + 1, row, created) - _finish(outcome) + _finish(outcome, block) except _BlockRollbackError: return @@ -466,3 +452,60 @@ class GeneratedCaseCountTest(SimpleTestCase): def test_the_generator_yields_every_valid_combination(self): """42 cases of one block, and 666 of two: 882 minus 216 that change the row after its delete.""" self.assertEqual(len(list(generated_cases())), 708) + + +class WriteScopeOnMainTest(TestCase): + """On main a write scope holds ``default`` alone, and the scope and its blocks add no query.""" + + def test_the_scope_pins_default_and_runs_no_query(self): + with self.assertNumQueries(0), write_scope() as aliases: + self.assertEqual(aliases, ("default",)) + + def test_an_unexpected_write_alias_raises_before_any_query(self): + with self.assertNumQueries(0), self.assertRaisesMessage(RuntimeError, "'default'") as raised: + with write_scope(expected_alias="schema_elsewhere"): + self.fail("the scope opened") + + self.assertIn("'schema_elsewhere'", str(raised.exception)) + + def test_a_nested_scope_joins_the_open_scope(self): + with write_scope() as outer, write_scope(expected_alias="default") as inner: + self.assertIs(inner, outer) + + def test_a_nested_scope_on_another_write_alias_raises_before_any_query(self): + with write_scope(), override_settings(DATABASE_ROUTERS=[WriteInterfacesTo("schema_elsewhere")]): + with self.assertNumQueries(0), self.assertRaisesMessage(RuntimeError, "'schema_elsewhere'"): + with write_scope(): + self.fail("the nested scope opened") + + def test_on_commit_outside_a_scope_raises(self): + with self.assertRaisesMessage(RuntimeError, "write scope"): + on_commit(lambda: None) + + def test_on_commit_registers_the_callback_itself_once(self): + def callback(): + pass + + with self.captureOnCommitCallbacks() as callbacks, write_scope(): + on_commit(callback) + + self.assertEqual(callbacks, [callback]) + + def test_a_block_marked_for_rollback_writes_nothing(self): + with atomic_with_events() as block: + self.assertEqual(block.aliases, ("default",)) + site = Site.objects.create(name="TxScope rollback", slug="txscope-rollback") + block.set_rollback() + + self.assertFalse(Site.objects.filter(pk=site.pk).exists()) + + def test_a_block_runs_the_statements_of_one_atomic_block(self): + with CaptureQueriesContext(connection) as plain, transaction.atomic(): + Site.objects.create(name="TxScope plain", slug="txscope-plain") + with CaptureQueriesContext(connection) as block, atomic_with_events(): + Site.objects.create(name="TxScope block", slug="txscope-block") + + def verbs(queries): + return [query["sql"].split()[0] for query in queries.captured_queries] + + self.assertEqual(verbs(block), verbs(plain)) diff --git a/netbox_interface_name_rules/tests/test_type_change_trigger.py b/netbox_interface_name_rules/tests/test_type_change_trigger.py index 15bb7938..a5e6acc9 100644 --- a/netbox_interface_name_rules/tests/test_type_change_trigger.py +++ b/netbox_interface_name_rules/tests/test_type_change_trigger.py @@ -23,24 +23,24 @@ from netbox_interface_name_rules.family import supports_module_moves from netbox_interface_name_rules.models import InterfaceNameRule from netbox_interface_name_rules.tests.committed_callbacks import run_the_reapply -from netbox_interface_name_rules.tests.test_bay_edit_trigger import BayEditTestCase, _flat_rule -from netbox_interface_name_rules.tests.test_module_move_trigger import ( +from netbox_interface_name_rules.tests.helpers import PLAIN_TYPE, REQUIRES_VC_POSITION_TOKEN +from netbox_interface_name_rules.tests.trigger_cases import ( FLAT, NAMING_READ, NO_RULE, - PLAIN_TYPE, REQUIRES_SUBTREE_MOVES, TAKEN, UNAVAILABLE, - _journal, - _MoveFixture, - _reapplied, + BayEditTestCase, + MoveFixture, + flat_rule, + journal, + reapplied, + reject_reads_of, ) -from netbox_interface_name_rules.tests.test_rename_triggers import _reject_reads_of -from netbox_interface_name_rules.tests.test_vc_drift import REQUIRES_VC_POSITION_TOKEN -class _CardFixture(_MoveFixture): +class _CardFixture(MoveFixture): """Two card types with two ports each, and an optic whose rule each card type scopes as its parent.""" @classmethod @@ -102,7 +102,7 @@ def test_a_nested_module_whose_rule_changes_is_renamed_from_the_names_the_old_ru self._change_type(card, self.second_card_type) self.assertEqual(self._names(optic), ["b-0/1"]) - self.assertEqual((_journal(card), _journal(optic)), ([], [])) + self.assertEqual((journal(card), journal(optic)), ([], [])) def test_a_nested_module_whose_rule_does_not_change_keeps_its_names_and_is_not_reported(self): card, second_port, optic = self._card_with_optic() @@ -113,7 +113,7 @@ def test_a_nested_module_whose_rule_does_not_change_keeps_its_names_and_is_not_r self._change_type(card, self.second_card_type) self.assertEqual((self._names(optic), self._names(fixed)), (["b-0/1"], ["et-1/0/2", "operator-name"])) - self.assertEqual((_journal(card), _journal(optic), _journal(fixed)), ([], [], [])) + self.assertEqual((journal(card), journal(optic), journal(fixed)), ([], [], [])) def test_a_module_two_levels_down_keeps_its_names_and_is_not_reported(self): card, second_port, optic = self._card_with_optic() @@ -129,7 +129,7 @@ def test_a_module_two_levels_down_keeps_its_names_and_is_not_reported(self): self._change_type(card, self.second_card_type) self.assertEqual((self._names(optic), self._names(deep)), (["b-0/1"], ["g-3", "operator-name"])) - self.assertEqual((_journal(card), _journal(sub_card), _journal(deep)), ([], [], [])) + self.assertEqual((journal(card), journal(sub_card), journal(deep)), ([], [], [])) def test_a_nested_module_left_without_a_rule_keeps_the_names_the_old_rule_gave_and_is_reported(self): card, _second_port, optic = self._card_with_optic() @@ -137,14 +137,14 @@ def test_a_nested_module_left_without_a_rule_keeps_the_names_the_old_rule_gave_a self._change_type(card, self._two_port_card("Bare Card")) self.assertEqual(self._names(optic), ["a-0/1"]) - (entry,) = _journal(card) + (entry,) = journal(card) self.assertEqual(entry.kind, JournalEntryKindChoices.KIND_WARNING) self.assertIn(f"`a-0/1`: {NO_RULE}", entry.comments) - self.assertEqual(_journal(optic), []) + self.assertEqual(journal(optic), []) def test_a_nested_flat_breakout_family_keeps_its_names_and_is_reported(self): flat_optic_type = self._module_type("Flat Optic", "{module}") - _flat_rule(flat_optic_type, "x-{slot}/{bay_position}:{channel}", parent_module_type=self.first_card_type) + flat_rule(flat_optic_type, "x-{slot}/{bay_position}:{channel}", parent_module_type=self.first_card_type) InterfaceNameRule.objects.create( module_type=flat_optic_type, parent_module_type=self.second_card_type, @@ -157,11 +157,11 @@ def test_a_nested_flat_breakout_family_keeps_its_names_and_is_reported(self): self._change_type(card, self.second_card_type) self.assertEqual(self._names(optic), ["x-0/1:0", "x-0/1:1"]) - (entry,) = _journal(card) + (entry,) = journal(card) self.assertEqual(entry.kind, JournalEntryKindChoices.KIND_WARNING) for name in ("x-0/1:0", "x-0/1:1"): self.assertIn(f"`{name}`: {FLAT}", entry.comments) - self.assertEqual(_journal(optic), []) + self.assertEqual(journal(optic), []) def test_the_subtree_reports_in_one_journal_entry_on_the_module_whose_type_changed(self): card, second_port, optic = self._card_with_optic() @@ -172,11 +172,11 @@ def test_the_subtree_reports_in_one_journal_entry_on_the_module_whose_type_chang self._change_type(card, self.second_card_type) self.assertEqual((self._names(optic), self._names(other_optic)), (["a-0/1"], ["a-0/2"])) - (entry,) = _journal(card) + (entry,) = journal(card) self.assertEqual(entry.kind, JournalEntryKindChoices.KIND_WARNING) self.assertIn(f"`a-0/1` to `b-0/1`: {TAKEN}", entry.comments) self.assertIn(f"`a-0/2` to `b-0/2`: {TAKEN}", entry.comments) - self.assertEqual((_journal(optic), _journal(other_optic)), ([], [])) + self.assertEqual((journal(optic), journal(other_optic)), ([], [])) class TypeChangeTransactionTest(TypeChangeTestCase): @@ -191,7 +191,7 @@ def test_a_type_change_then_a_bay_edit_rename_the_nested_module_from_the_names_b self._save_edit(bay, position="2") self.assertEqual(self._names(optic), ["b-2/1"]) - self.assertEqual((_journal(card), _journal(optic)), ([], [])) + self.assertEqual((journal(card), journal(optic)), ([], [])) @skipUnless(supports_module_moves(), REQUIRES_SUBTREE_MOVES) def test_a_type_change_then_a_move_rename_the_nested_module_from_the_names_before_the_type_change(self): @@ -202,7 +202,7 @@ def test_a_type_change_then_a_move_rename_the_nested_module_from_the_names_befor self._save_move(card, self._bay(self.device, "Bay 2")) self.assertEqual(self._names(optic), ["b-2/1"]) - self.assertEqual((_journal(card), _journal(optic)), ([], [])) + self.assertEqual((journal(card), journal(optic)), ([], [])) class ChassisPositionTypeChangeTest(TypeChangeTestCase): @@ -245,8 +245,8 @@ def _assert_the_nested_module_is_renamed_once(self, chassis_first): reapplies = self._retype_with_the_chassis_change(card, self.second_card_type, chassis_first) self.assertEqual((self._names(optic), self._names(other)), (["y30/1"], ["et-3/0/10"])) - self.assertEqual(_reapplied(reapplies), sorted((card.pk, optic.pk, other.pk))) - self.assertEqual((_journal(card), _journal(optic), _journal(self.device)), ([], [], [])) + self.assertEqual(reapplied(reapplies), sorted((card.pk, optic.pk, other.pk))) + self.assertEqual((journal(card), journal(optic), journal(self.device)), ([], [], [])) def test_a_type_change_then_a_chassis_position_change_rename_the_nested_module_once(self): self._assert_the_nested_module_is_renamed_once(chassis_first=False) @@ -262,8 +262,8 @@ def _assert_an_unchanged_nested_rule_is_left_to_the_chassis_position_change(self reapplies = self._retype_with_the_chassis_change(card, self.second_card_type, chassis_first) self.assertEqual((self._names(optic), self._names(fixed)), (["b-0/1"], ["et-3/0/2"])) - self.assertEqual(_reapplied(reapplies), sorted((card.pk, optic.pk, fixed.pk))) - self.assertEqual((_journal(card), _journal(fixed), _journal(self.device)), ([], [], [])) + self.assertEqual(reapplied(reapplies), sorted((card.pk, optic.pk, fixed.pk))) + self.assertEqual((journal(card), journal(fixed), journal(self.device)), ([], [], [])) def test_a_type_change_then_a_chassis_position_change_rename_a_nested_module_whose_rule_does_not_change(self): self._assert_an_unchanged_nested_rule_is_left_to_the_chassis_position_change(chassis_first=False) @@ -280,9 +280,9 @@ def _assert_leaving_with_a_type_change_renames_nothing_and_reports_once(self, le ) self.assertEqual((self._names(module), self._names(other)), (["et-1/0/0"], ["et-1/0/10"])) - self.assertEqual(_reapplied(reapplies), sorted((module.pk, other.pk))) - (module_entry,) = _journal(module) - (device_entry,) = _journal(self.device) + self.assertEqual(reapplied(reapplies), sorted((module.pk, other.pk))) + (module_entry,) = journal(module) + (device_entry,) = journal(self.device) self.assertEqual(module_entry.comments.count(f"`et-1/0/0`: {UNAVAILABLE}"), 1) self.assertIn(f"`et-1/0/10`: {UNAVAILABLE}", device_entry.comments) self.assertNotIn("`et-1/0/0`", device_entry.comments) @@ -299,8 +299,8 @@ def _assert_a_raw_name_with_adjacent_tokens_is_renamed_by_the_forced_reapply(sel reapplies = self._retype_with_the_chassis_change(module, self.adjacent_type, chassis_first) - self.assertEqual((self._names(module), _reapplied(reapplies)), (["ge-3/0"], [module.pk])) - self.assertEqual((_journal(module), _journal(self.device)), ([], [])) + self.assertEqual((self._names(module), reapplied(reapplies)), (["ge-3/0"], [module.pk])) + self.assertEqual((journal(module), journal(self.device)), ([], [])) @skipUnless(supports_vc_position_token(), REQUIRES_VC_POSITION_TOKEN) def test_a_type_change_then_a_chassis_position_change_rename_a_raw_name_with_adjacent_tokens(self): @@ -316,10 +316,10 @@ def _assert_a_collision_is_reported_once(self, chassis_first): reapplies = self._retype_with_the_chassis_change(module, self.other_plain_type, chassis_first) - self.assertEqual((self._names(module), _reapplied(reapplies)), (["et-1/0/0"], [module.pk])) - (entry,) = _journal(module) + self.assertEqual((self._names(module), reapplied(reapplies)), (["et-1/0/0"], [module.pk])) + (entry,) = journal(module) self.assertEqual(entry.comments.count(f"`et-1/0/0` to `ge-3/0/0`: {TAKEN}"), 1) - self.assertEqual(_journal(self.device), []) + self.assertEqual(journal(self.device), []) def test_a_type_change_then_a_chassis_position_change_report_a_collision_once(self): self._assert_a_collision_is_reported_once(chassis_first=False) @@ -340,18 +340,18 @@ def test_a_rule_read_failure_is_reported_for_each_module_and_the_later_modules_s failure = f"injected {InterfaceNameRule._meta.db_table} read failure" with ( - connection.execute_wrapper(_reject_reads_of(InterfaceNameRule._meta.db_table)), + connection.execute_wrapper(reject_reads_of(InterfaceNameRule._meta.db_table)), self.assertLogs("netbox_interface_name_rules", "ERROR"), ): run_the_reapply(callbacks) - (card_entry,) = _journal(card) + (card_entry,) = journal(card) self.assertEqual(card_entry.kind, JournalEntryKindChoices.KIND_DANGER) self.assertEqual(card_entry.comments.count(failure), 2) - (other_entry,) = _journal(other) + (other_entry,) = journal(other) self.assertEqual(other_entry.kind, JournalEntryKindChoices.KIND_DANGER) self.assertEqual(other_entry.comments.count(failure), 1) - self.assertEqual((self._names(optic), self._names(other), _journal(optic)), (["a-0/1"], ["et-1/0/1"], [])) + self.assertEqual((self._names(optic), self._names(other), journal(optic)), (["a-0/1"], ["et-1/0/1"], [])) class TypeChangeCostTest(TypeChangeTestCase): @@ -378,7 +378,7 @@ def test_a_type_change_reads_no_nested_naming_without_an_enabled_rule_scoped_to_ self.assertEqual(reads, (0, 1)) self.assertEqual(self._names(optic), ["u-0/1"]) - self.assertEqual(_journal(card), []) + self.assertEqual(journal(card), []) def test_a_type_change_reads_the_nested_naming_when_an_enabled_rule_is_scoped_to_the_new_type(self): card, _second_port, optic = self._card_with_optic() @@ -418,4 +418,4 @@ def test_patching_the_module_type_renames_the_nested_module(self): self.assertEqual(response.status_code, status.HTTP_200_OK, response.data) self.assertEqual(self._names(optic), ["b-0/1"]) - self.assertEqual(_journal(card), []) + self.assertEqual(journal(card), []) diff --git a/netbox_interface_name_rules/tests/test_vc_drift.py b/netbox_interface_name_rules/tests/test_vc_drift.py index 0b3f5d68..8c72f14b 100644 --- a/netbox_interface_name_rules/tests/test_vc_drift.py +++ b/netbox_interface_name_rules/tests/test_vc_drift.py @@ -48,7 +48,6 @@ from extras.models import JournalEntry from netbox_interface_name_rules import engine -from netbox_interface_name_rules.choices import BreakoutModeChoices from netbox_interface_name_rules.engine import ( apply_interface_name_rules, apply_rule_to_existing, @@ -63,33 +62,23 @@ from netbox_interface_name_rules.models import InterfaceNameRule from netbox_interface_name_rules.naming import build_variables from netbox_interface_name_rules.rename_triggers import ModuleTrigger, reapply -from netbox_interface_name_rules.tests.out_of_band import rename_out_of_band -from netbox_interface_name_rules.tests.test_channelization import ( +from netbox_interface_name_rules.tests.helpers import ( CHANNEL_TYPE, + CHANNELIZED, + FLAT, PARENT_TYPE, PLAIN_TYPE, PLUGIN_LOGGER, REQUIRES_CHANNELIZATION, - ChannelizationTestCase, - _build_device, + REQUIRES_VC_POSITION_TOKEN, + VcDriftTestCase, + build_device, + token_module_type, ) +from netbox_interface_name_rules.tests.out_of_band import rename_out_of_band -FLAT = BreakoutModeChoices.FLAT -CHANNELIZED = BreakoutModeChoices.CHANNELIZED CLAIM_LOGGER = "netbox_interface_name_rules.family.raw_bases" -# Every fixture spelling a name NetBox resolved from the token needs the release that resolves it: -# on 4.5 and older the token stays literal in the interface name and the drift cannot even occur. -REQUIRES_VC_POSITION_TOKEN = "requires a NetBox that resolves {vc_position} in template names (4.6+)" # noqa: S105 - Skip reason, not a credential. - - -def _token_module_type(manufacturer, model, *template_names, iface_type=PLAIN_TYPE): - """Create a ModuleType whose interface templates are named *template_names*, in order.""" - module_type = ModuleType.objects.create(manufacturer=manufacturer, model=model, part_number=model) - for name in template_names: - InterfaceTemplate.objects.create(module_type=module_type, name=name, type=iface_type) - return module_type - def _raw_name_patterns(module): """Return the structural raw-name matchers of *module*'s templates (see the module docstring).""" @@ -120,36 +109,6 @@ def _without_vc_position_re(): return mock.patch.dict(sys.modules, {"dcim.constants": _ConstantsWithoutVcToken(dcim.constants)}) -class VcDriftTestCase(ChannelizationTestCase): - """VC transitions go through a real ``Device.save()`` so the plugin's signals do the scheduling.""" - - def _save_vc_state(self, device, virtual_chassis, position): - with self.captureOnCommitCallbacks(execute=True): - device.virtual_chassis = virtual_chassis - device.vc_position = position - device.save() - - def _join(self, vc, position, device=None): - """Add *device* to *vc* at *position* — the join direction (fallback → position).""" - self._save_vc_state(device or self.device, vc, position) - - def _renumber(self, position, device=None): - """Move *device* to another position inside its VC — the renumber direction (P → Q).""" - device = device or self.device - self._save_vc_state(device, device.virtual_chassis, position) - - def _leave(self, device=None): - """Remove *device* from its VC — the leave direction (position → fallback).""" - self._save_vc_state(device or self.device, None, None) - - def _install_on(self, device, module_type, position): - """Install a module of *module_type* into *device*'s bay at *position*, rules and all.""" - bay = ModuleBay.objects.get(device=device, name=f"Bay {position}") - with self.captureOnCommitCallbacks(execute=True): - module = Module.objects.create(device=device, module_bay=bay, module_type=module_type) - return module, bay - - # --------------------------------------------------------------------------- # Join: the device was standalone when NetBox named the interfaces # --------------------------------------------------------------------------- @@ -161,8 +120,8 @@ class VcPositionJoinDriftTest(VcDriftTestCase): @classmethod def setUpTestData(cls): - manufacturer, cls.device = _build_device("VcJoin", ["3", "4"]) - cls.module_type = _token_module_type(manufacturer, "VcJoin-QSFP", "xe-{vc_position:0}/0/{module}") + manufacturer, cls.device = build_device("VcJoin", ["3", "4"]) + cls.module_type = token_module_type(manufacturer, "VcJoin-QSFP", "xe-{vc_position:0}/0/{module}") def test_joining_a_vc_renames_a_flat_channel_family_named_with_the_fallback(self): """The worst case from the issue: under force, a breakout rule matches its family by raw name.""" @@ -206,11 +165,11 @@ class VcPositionRenumberDriftTest(VcDriftTestCase): @classmethod def setUpTestData(cls): - manufacturer, cls.device = _build_device( + manufacturer, cls.device = build_device( "VcRenum", ["3", "5"], virtual_chassis=VirtualChassis.objects.create(name="vcrenum-vc"), vc_position=1 ) - cls.module_type = _token_module_type(manufacturer, "VcRenum-QSFP", "xe-{vc_position:0}/0/{module}") - cls.simple_type = _token_module_type(manufacturer, "VcRenum-SFP", "xe-{vc_position:0}/0/{module}") + cls.module_type = token_module_type(manufacturer, "VcRenum-QSFP", "xe-{vc_position:0}/0/{module}") + cls.simple_type = token_module_type(manufacturer, "VcRenum-SFP", "xe-{vc_position:0}/0/{module}") def test_renumbering_renames_a_family_named_at_an_earlier_position(self): """The token sits in a middle path segment, so the drifted name is structural, not a suffix.""" @@ -248,10 +207,10 @@ class VcPositionForceBaseMatchingTest(VcDriftTestCase): @classmethod def setUpTestData(cls): - manufacturer, cls.device = _build_device( + manufacturer, cls.device = build_device( "VcForm", ["5", "7"], virtual_chassis=VirtualChassis.objects.create(name="vcform-vc"), vc_position=1 ) - cls.module_type = _token_module_type(manufacturer, "VcForm-QSFP", "{vc_position}-{module}") + cls.module_type = token_module_type(manufacturer, "VcForm-QSFP", "{vc_position}-{module}") def _breakout_rule(self, name_template): return InterfaceNameRule.objects.create( @@ -289,10 +248,10 @@ class VcPositionLeaveDriftTest(VcDriftTestCase): @classmethod def setUpTestData(cls): - manufacturer, cls.device = _build_device( + manufacturer, cls.device = build_device( "VcLeave", ["3", "4", "5"], virtual_chassis=VirtualChassis.objects.create(name="vcleave-vc"), vc_position=2 ) - cls.module_type = _token_module_type(manufacturer, "VcLeave-SFP", "xe-{vc_position:0}/0/{module}") + cls.module_type = token_module_type(manufacturer, "VcLeave-SFP", "xe-{vc_position:0}/0/{module}") def test_leaving_a_vc_renames_nothing(self): """Deliberate: what the interfaces are called off a VC is an operator decision, even for a rule @@ -334,16 +293,14 @@ class VcPositionAmbiguityTest(VcDriftTestCase): @classmethod def setUpTestData(cls): - manufacturer, cls.device = _build_device( + manufacturer, cls.device = build_device( "VcAmb", ["3", "4"], virtual_chassis=VirtualChassis.objects.create(name="vcamb-vc"), vc_position=1 ) # One token template plus a plain one whose interface an earlier rename moved onto the # token template's fallback variant. - cls.decoy_type = _token_module_type( - manufacturer, "VcAmb-QSFP", "xe-{vc_position:0}/0/{module}", "mgmt-{module}" - ) + cls.decoy_type = token_module_type(manufacturer, "VcAmb-QSFP", "xe-{vc_position:0}/0/{module}", "mgmt-{module}") # Two token templates whose matchers overlap on 'xe-1/0/4'. - cls.overlap_type = _token_module_type( + cls.overlap_type = token_module_type( manufacturer, "VcAmb-SFP", "xe-{vc_position}/0/{module}", "xe-1/{vc_position}/{module}" ) @@ -506,10 +463,10 @@ class VcPositionOneClaimPassTest(VcDriftTestCase): @classmethod def setUpTestData(cls): - manufacturer, cls.device = _build_device( + manufacturer, cls.device = build_device( "VcPass", ["0"], virtual_chassis=VirtualChassis.objects.create(name="vcpass-vc"), vc_position=1 ) - cls.module_type = _token_module_type(manufacturer, "VcPass-QSFP", "{vc_position}/{module}") + cls.module_type = token_module_type(manufacturer, "VcPass-QSFP", "{vc_position}/{module}") def _flat_rule(self): return InterfaceNameRule.objects.create( @@ -602,11 +559,11 @@ class VcPositionAdjacentTokenTest(VcDriftTestCase): @classmethod def setUpTestData(cls): - manufacturer, cls.device = _build_device( + manufacturer, cls.device = build_device( "VcAdj", ["3", "4"], virtual_chassis=VirtualChassis.objects.create(name="vcadj-vc"), vc_position=1 ) - cls.adjacent_type = _token_module_type(manufacturer, "VcAdj-QSFP", "xe-{vc_position}{vc_position}/0/{module}") - cls.separated_type = _token_module_type(manufacturer, "VcAdj-SFP", "xe-{vc_position}/{vc_position}/{module}") + cls.adjacent_type = token_module_type(manufacturer, "VcAdj-QSFP", "xe-{vc_position}{vc_position}/0/{module}") + cls.separated_type = token_module_type(manufacturer, "VcAdj-SFP", "xe-{vc_position}/{vc_position}/{module}") @skipUnless(supports_vc_position_token(), REQUIRES_VC_POSITION_TOKEN) def test_adjacent_tokens_build_no_matcher_at_all(self): @@ -635,10 +592,10 @@ class VcPositionResolutionTest(VcDriftTestCase): @classmethod def setUpTestData(cls): - manufacturer, cls.device = _build_device( + manufacturer, cls.device = build_device( "VcRes", ["3"], virtual_chassis=VirtualChassis.objects.create(name="vcres-vc"), vc_position=1 ) - cls.module_type = _token_module_type( + cls.module_type = token_module_type( manufacturer, "VcRes-QSFP", "xe-{vc_position}/{module}", @@ -698,8 +655,8 @@ class VcPositionNoTokenControlTest(VcDriftTestCase): @classmethod def setUpTestData(cls): - manufacturer, cls.device = _build_device("VcNoTok", ["3", "4"]) - cls.module_type = _token_module_type(manufacturer, "VcNoTok-QSFP", "{module}") + manufacturer, cls.device = build_device("VcNoTok", ["3", "4"]) + cls.module_type = token_module_type(manufacturer, "VcNoTok-QSFP", "{module}") @skipUnless(supports_vc_position_token(), REQUIRES_VC_POSITION_TOKEN) def test_a_tokenless_module_type_builds_no_matchers(self): @@ -735,10 +692,10 @@ class VcPositionLegacyNetboxTest(VcDriftTestCase): @classmethod def setUpTestData(cls): - manufacturer, cls.device = _build_device( + manufacturer, cls.device = build_device( "VcLegacy", ["3"], virtual_chassis=VirtualChassis.objects.create(name="vclegacy-vc"), vc_position=1 ) - cls.module_type = _token_module_type(manufacturer, "VcLegacy-SFP", "xe-{vc_position:0}/0/{module}") + cls.module_type = token_module_type(manufacturer, "VcLegacy-SFP", "xe-{vc_position:0}/0/{module}") def test_the_feature_check_is_false_without_the_constant(self): """Probed from ``dcim.constants``, lazily — an upstream removal must flip the check, not crash.""" @@ -777,14 +734,14 @@ class VcPositionNestedBayTest(VcDriftTestCase): @classmethod def setUpTestData(cls): - manufacturer, cls.device = _build_device( + manufacturer, cls.device = build_device( "VcNest", ["2"], virtual_chassis=VirtualChassis.objects.create(name="vcnest-vc"), vc_position=1 ) cls.chassis_type = ModuleType.objects.create( manufacturer=manufacturer, model="VcNest-Chassis", part_number="VcNest-Chassis" ) ModuleBayTemplate.objects.create(module_type=cls.chassis_type, name="LC Bay", position="1") - cls.leaf_type = _token_module_type(manufacturer, "VcNest-LEAF", "xe-{vc_position:0}/{module}") + cls.leaf_type = token_module_type(manufacturer, "VcNest-LEAF", "xe-{vc_position:0}/{module}") def _install_leaf(self): """Install the chassis in the device bay and a leaf module in the chassis' own bay.""" @@ -833,10 +790,10 @@ class VcPositionPredictionTest(VcDriftTestCase): @classmethod def setUpTestData(cls): - manufacturer, cls.device = _build_device( + manufacturer, cls.device = build_device( "VcPred", ["3"], virtual_chassis=VirtualChassis.objects.create(name="vcpred-vc"), vc_position=2 ) - cls.module_type = _token_module_type(manufacturer, "VcPred-SFP", "xe-{vc_position:0}/0/{module}") + cls.module_type = token_module_type(manufacturer, "VcPred-SFP", "xe-{vc_position:0}/0/{module}") def test_names_resolved_at_call_time_still_predict_correctly(self): """The documented precondition: the caller resolves the names, so same-instant input is exact.""" @@ -865,7 +822,7 @@ class VcPositionAsymmetricFamilyTest(VcDriftTestCase): @classmethod def setUpTestData(cls): - manufacturer, cls.device = _build_device( + manufacturer, cls.device = build_device( "VcAsym", ["3"], virtual_chassis=VirtualChassis.objects.create(name="vcasym-vc"), vc_position=1 ) cls.module_type = ModuleType.objects.create( @@ -908,18 +865,18 @@ class VcPositionConversionRecoveryTest(VcDriftTestCase): @classmethod def setUpTestData(cls): - manufacturer, cls.device = _build_device( + manufacturer, cls.device = build_device( "VcConv", ["3", "4", "5", "6", "7"], virtual_chassis=VirtualChassis.objects.create(name="vcconv-vc"), vc_position=1, ) cls.manufacturer = manufacturer - cls.wrap_type = _token_module_type(manufacturer, "VcConv-WRAP", "xe-{vc_position:0}/0/{module}") - cls.twice_type = _token_module_type(manufacturer, "VcConv-TWICE", "xe-{vc_position:0}/0/{module}") - cls.arith_type = _token_module_type(manufacturer, "VcConv-ARITH", "{vc_position}{module}") - cls.free_type = _token_module_type(manufacturer, "VcConv-FREE", "xe-{vc_position:0}/0/{module}") - cls.two_base_type = _token_module_type( + cls.wrap_type = token_module_type(manufacturer, "VcConv-WRAP", "xe-{vc_position:0}/0/{module}") + cls.twice_type = token_module_type(manufacturer, "VcConv-TWICE", "xe-{vc_position:0}/0/{module}") + cls.arith_type = token_module_type(manufacturer, "VcConv-ARITH", "{vc_position}{module}") + cls.free_type = token_module_type(manufacturer, "VcConv-FREE", "xe-{vc_position:0}/0/{module}") + cls.two_base_type = token_module_type( manufacturer, "VcConv-TWOBASE", "xe-{vc_position:0}/0/{module}", "xe-{vc_position:9}/0/{module}" ) @@ -1012,7 +969,7 @@ def test_a_family_matcher_that_captures_two_bases_is_not_offered(self): def test_an_unrelated_family_survives_overlapping_historical_claims(self): """A multi-base claim rejects its shared base but leaves an unrelated family available.""" - module_type = _token_module_type( + module_type = token_module_type( self.manufacturer, "VcConv-OVERLAP", "xe-{vc_position}/0/{module}", @@ -1044,7 +1001,7 @@ def test_an_unrelated_family_survives_overlapping_historical_claims(self): def test_a_family_a_current_and_a_historical_base_both_spell_is_not_offered(self): """The rule gives the family to one template now and to the other at position 1: neither converts it.""" - module_type = _token_module_type( + module_type = token_module_type( self.manufacturer, "VcConv-CURRENT-FIRST", "xe-{vc_position:0}/0/{module}", @@ -1081,7 +1038,7 @@ def test_a_rule_without_a_base_is_identified_after_a_renumber(self): def test_a_family_a_current_and_a_historical_base_both_spell_is_not_planned(self): """One claim over every form: a name the rule gives two templates is neither template's.""" - module_type = _token_module_type( + module_type = token_module_type( self.manufacturer, "VcConv-CURRENT-OVERLAP", "xe-{vc_position:0}/0/{module}", diff --git a/netbox_interface_name_rules/tests/test_views.py b/netbox_interface_name_rules/tests/test_views.py index 91db1b63..1547141d 100644 --- a/netbox_interface_name_rules/tests/test_views.py +++ b/netbox_interface_name_rules/tests/test_views.py @@ -4,8 +4,9 @@ import re from html.parser import HTMLParser -from unittest.mock import ANY, MagicMock, patch +from unittest.mock import patch +from core.models import Job from dcim.models import ( DeviceType, Interface, @@ -20,18 +21,18 @@ from django.urls import path, re_path, reverse from rest_framework.test import APIClient +from netbox_interface_name_rules.jobs import ApplyRuleJob, ConvertFlatFamiliesJob from netbox_interface_name_rules.models import InterfaceNameRule from netbox_interface_name_rules.name_template import TEMPLATE_VARIABLES, NamingContext, variables_for_context from netbox_interface_name_rules.template_variable_reference import ( rule_tester_variable_rows, variable_reference_rows, ) -from netbox_interface_name_rules.tests.helpers import make_device +from netbox_interface_name_rules.tests.helpers import TEST_PASSWORD, make_device, queued_job from netbox_interface_name_rules.views import RuleTestView User = get_user_model() -TEST_PASSWORD = "testpass123" # noqa: S105 - Test credential only. # The preview variables and override fields the conftest guard is expected to know about. _VAR_FIELDS = frozenset({"slot", "bay_position", "parent_bay_position", "base", "vc_position"}) @@ -1429,17 +1430,19 @@ def test_convert_failure_is_reported_rather_than_raised(self): self.assertEqual(response.status_code, 302) self.assertTrue(any(m.level == ERROR for m in get_messages(response.wsgi_request))) - def test_convert_background_enqueues_the_conversion_job(self): - """The batch runs on the worker, the same way a large apply does.""" - mock_job = MagicMock() - mock_job.pk = 43 - with patch( - "netbox_interface_name_rules.jobs.ConvertFlatFamiliesJob.enqueue", return_value=mock_job - ) as mock_enqueue: + def test_convert_background_enqueues_the_conversion_job_on_main(self): + """The batch runs on the worker, the same way a large apply does, on main when no branch is active.""" + with self.captureOnCommitCallbacks(execute=True): response = self.client.post(self._url(), {"action": "convert_background"}) self.assertEqual(response.status_code, 302) - mock_enqueue.assert_called_once_with(name=ANY, user=self.superuser, rule_id=self.rule.pk) + job = Job.objects.get(name=f"Convert flat families: {self.rule}", user=self.superuser) + queued = queued_job(self, job) + self.assertEqual(queued.func, ConvertFlatFamiliesJob.handle) + self.assertEqual( + queued.kwargs, + {"job": job, "rule_id": self.rule.pk, "branch_schema_id": None, "expected_alias": "default"}, + ) class RuleToggleNoPermissionNonAjaxTest(ViewTestBase2): @@ -1619,28 +1622,29 @@ def test_get_re_error_shows_error_message(self): class RuleApplyDetailViewBackgroundJobSuccessTest(ViewTestBase2): """Test RuleApplyDetailView.post with background action that succeeds.""" - def test_post_background_success_shows_success_message(self): - """POST background action with successful enqueue shows success message (line 400). - - Asserts ApplyRuleJob.enqueue is called once and the success message - contains the job pk (42) confirming the enqueued job id is reported. - """ - from django.contrib.messages import SUCCESS, get_messages + def test_post_background_enqueues_the_job_on_main_and_reports_its_id(self): + """The job stores the rule, no branch and the write alias of main, and the operator gets its ID.""" + from django.contrib.messages import get_messages url = reverse( "plugins:netbox_interface_name_rules:interfacenamerule_apply_detail", kwargs={"pk": self.rule.pk}, ) - mock_job = MagicMock() - mock_job.pk = 42 - with patch("netbox_interface_name_rules.jobs.ApplyRuleJob.enqueue", return_value=mock_job) as mock_enq: + with self.captureOnCommitCallbacks(execute=True): response = self.client.post(url, {"action": "background"}) + self.assertEqual(response.status_code, 302) - mock_enq.assert_called_once_with(name=ANY, user=self.superuser, rule_id=self.rule.pk) - msgs = list(get_messages(response.wsgi_request)) - success_msgs = [m for m in msgs if m.level == SUCCESS] - self.assertTrue(success_msgs, "Expected a success-level message but none found") - self.assertTrue(any("42" in str(m) for m in success_msgs)) + job = Job.objects.get(name=f"Apply rule: {self.rule}", user=self.superuser) + queued = queued_job(self, job) + self.assertEqual(queued.func, ApplyRuleJob.handle) + self.assertEqual( + queued.kwargs, + {"job": job, "rule_id": self.rule.pk, "branch_schema_id": None, "expected_alias": "default"}, + ) + self.assertEqual( + [(message.level_tag, str(message)) for message in get_messages(response.wsgi_request)], + [("success", f"Background job enqueued (job #{job.pk}). Check Core → Jobs for status.")], + ) # --------------------------------------------------------------------------- diff --git a/netbox_interface_name_rules/tests/trigger_cases.py b/netbox_interface_name_rules/tests/trigger_cases.py new file mode 100644 index 00000000..9f3af9fc --- /dev/null +++ b/netbox_interface_name_rules/tests/trigger_cases.py @@ -0,0 +1,292 @@ +# SPDX-License-Identifier: Apache-2.0 +# Copyright (C) 2025 Marcin Zieba +"""Fixtures and probes that several rename-trigger test modules share. + +``MoveFixture`` builds two device types with the same bays and three devices in two virtual chassis. +``ModuleMoveTestCase`` and ``BayEditTestCase`` install, move and edit through real saves. The probes +count the reapplies, read the journal and inject database failures. +""" + +import re +from contextlib import contextmanager +from typing import NamedTuple +from unittest.mock import patch + +from dcim.models import Interface, InterfaceTemplate, Module, ModuleBay, ModuleBayTemplate, Platform, VirtualChassis +from django.contrib.contenttypes.models import ContentType +from django.db import DatabaseError, IntegrityError, connection, transaction +from django.test import TestCase +from django.test.utils import CaptureQueriesContext +from extras.models import JournalEntry + +from netbox_interface_name_rules import engine +from netbox_interface_name_rules.choices import BreakoutModeChoices +from netbox_interface_name_rules.models import InterfaceNameRule +from netbox_interface_name_rules.tests.helpers import ( + PLAIN_TYPE, + make_device, + make_device_type, + make_manufacturer, + make_module_type, + make_placement, + slug_for, +) + +BAYS = (("Bay 0", "0"), ("Bay 1", "1"), ("Bay 2", "2"), ("Bay 10", "10")) +REQUIRES_SUBTREE_MOVES = "requires a NetBox that moves a module's nested bays with it (4.7+)" +FLAT = "a flat breakout family is not renamed after a move, a bay edit or a parent module type change" +NO_RULE = "no rule matches the module after the change" +TAKEN = "target name is already in use" +UNAVAILABLE = "{vc_position} is not available on this device" +NAMING_READ = re.compile(r'SELECT .* FROM "dcim_module" .*"dcim_platform"') + + +class ChassisRule(NamedTuple): + """One rule shape that the tests of a chassis change with a module change cover.""" + + model: str + name_template: str + + +# A plain rule, a rule that reads {base}, and a rule with the virtual-chassis position in arithmetic. +CHASSIS_RULES = ( + ChassisRule("Plain", "et-{vc_position}/{slot}/{bay_position}"), + ChassisRule("Base", "p{base}-{vc_position}/{slot}"), + ChassisRule("Arithmetic", "x{{vc_position} * 10 + {slot_num}}/{bay_position}"), +) + + +def reject_interface_updates(execute, sql, params, many, context): + """Fail each update of an interface, as an execute wrapper.""" + if sql.lstrip().startswith('UPDATE "dcim_interface"'): + raise IntegrityError("injected reapply failure") + return execute(sql, params, many, context) + + +def fail_the_naming_read(execute, sql, params, many, context): + """Replace the subtree naming read by SQL that PostgreSQL rejects, as an execute wrapper.""" + if NAMING_READ.match(sql): + return execute("SELECT 1/0", None, many, context) + return execute(sql, params, many, context) + + +def journal(instance): + """Return the journal entries on *instance*, oldest first.""" + return list( + JournalEntry.objects.filter( + assigned_object_type=ContentType.objects.get_for_model(instance), assigned_object_id=instance.pk + ).order_by("pk") + ) + + +@contextmanager +def module_reapplies(): + """Count the module reapplies; each call still runs the real function.""" + with patch.object(engine, "module_rule_outcomes", wraps=engine.module_rule_outcomes) as spy: + yield spy + + +def reapplied(spy): + """Return the primary key of the module of each reapply that *spy* recorded, sorted.""" + return sorted(call.args[0].pk for call in spy.call_args_list) + + +@contextmanager +def naming_reads(): + """Record the queries and the result of each subtree naming read; each call still runs the real function.""" + reads = [] + real = engine.read_subtree_naming + + def read(module_pk): + with CaptureQueriesContext(connection) as queries: + naming = real(module_pk) + reads.append((queries.captured_queries, naming)) + return naming + + with patch.object(engine, "read_subtree_naming", read): + yield reads + + +def give_the_next_module_id(pk): + """Make the database assign *pk* to the next module that NetBox creates.""" + # NetBox before 4.7 creates no components for a module saved with an explicit pk. + with connection.cursor() as cursor: + cursor.execute( + "SELECT setval(pg_get_serial_sequence(%s, %s), %s, false)", + [Module._meta.db_table, Module._meta.pk.column, pk], + ) + + +@contextmanager +def previous_state_read_fails(statement): + """Replace the previous-state read that *statement* matches by SQL that PostgreSQL rejects.""" + replaced = [] + + def divide_by_zero(execute, sql, params, many, context): + if statement.match(sql): + replaced.append(sql) + return execute("SELECT 1/0", None, many, context) + return execute(sql, params, many, context) + + with connection.execute_wrapper(divide_by_zero): + yield replaced + + +def reject_reads_of(table): + """Return an execute wrapper that fails every read of *table*.""" + + def reject(execute, sql, params, many, context): + if sql.lstrip().startswith("SELECT") and (f'FROM "{table}"' in sql or f'JOIN "{table}"' in sql): + raise DatabaseError(f"injected {table} read failure") + return execute(sql, params, many, context) + + return reject + + +class MoveFixture: + """Two device types with the same bays, and three devices in two virtual chassis. + + ``device`` and ``peer`` are virtual-chassis positions 1 and 2 of one chassis; ``remote`` has + another device type and platform, at position 5 of another chassis. + """ + + @classmethod + def build(cls, prefix): + """Create the fixture objects and return them as class attributes of *cls*.""" + cls.prefix = prefix + cls.manufacturer = make_manufacturer(prefix) + cls.device_type = make_device_type(cls.manufacturer, prefix) + cls.other_device_type = make_device_type(cls.manufacturer, f"{prefix} Other") + for device_type in (cls.device_type, cls.other_device_type): + for name, position in BAYS: + ModuleBayTemplate.objects.create(device_type=device_type, name=name, position=position) + cls.platform = Platform.objects.create(name=f"{prefix} OS", slug=slug_for(prefix, "os")) + cls.other_platform = Platform.objects.create(name=f"{prefix} Other OS", slug=slug_for(prefix, "other-os")) + placement = make_placement(prefix) + chassis = VirtualChassis.objects.create(name=f"{prefix} VC") + remote_chassis = VirtualChassis.objects.create(name=f"{prefix} Remote VC") + cls.device = cls._device(placement, "01", cls.device_type, cls.platform, chassis, 1) + cls.peer = cls._device(placement, "02", cls.device_type, cls.other_platform, chassis, 2) + cls.remote = cls._device(placement, "03", cls.other_device_type, cls.other_platform, remote_chassis, 5) + + @classmethod + def _device(cls, placement, suffix, device_type, platform, chassis, position): + return make_device( + cls.prefix, + device_type, + placement, + name=slug_for(cls.prefix, suffix), + platform=platform, + virtual_chassis=chassis, + vc_position=position, + ) + + @classmethod + def _module_type(cls, model, *templates): + module_type = make_module_type(cls.manufacturer, model, model=f"{cls.prefix} {model}") + for template in templates: + InterfaceTemplate.objects.create(module_type=module_type, name=template, type=PLAIN_TYPE) + return module_type + + @classmethod + def _card_type(cls, model, bay_position): + """Return a module type that holds one nested bay at *bay_position*, and no interfaces.""" + card_type = make_module_type(cls.manufacturer, model, model=f"{cls.prefix} {model}") + ModuleBayTemplate.objects.create(module_type=card_type, name="Port", position=bay_position) + return card_type + + @staticmethod + def _bay(device, name="Bay 0"): + return ModuleBay.objects.get(device=device, module__isnull=True, name=name) + + @staticmethod + def _names(module): + return sorted(Interface.objects.filter(module=module).values_list("name", flat=True)) + + +class ModuleMoveTestCase(MoveFixture, TestCase): + """Install and move modules through real saves, with the committed callbacks run.""" + + @classmethod + def setUpTestData(cls): + cls.build(cls.__name__) + + def _install(self, module_type, bay): + with self.captureOnCommitCallbacks(execute=True): + return Module.objects.create(device=bay.device, module_bay=bay, module_type=module_type) + + @staticmethod + def _save_move(module, bay): + module.device = bay.device + module.module_bay = bay + module.save() + + def _move(self, module, bay): + with self.captureOnCommitCallbacks(execute=True): + self._save_move(module, bay) + + def _install_card(self, card_type, bay): + """Install *card_type* in *bay* and return it with its nested bay.""" + card = self._install(card_type, bay) + return card, ModuleBay.objects.get(module=card) + + def _change_the_chassis_position(self, position=3): + self.device.vc_position = position + self.device.save() + + def _leave_the_chassis(self): + self.device.virtual_chassis = None + self.device.vc_position = None + self.device.save() + + def _join_the_chassis(self, chassis): + self.device.virtual_chassis = chassis + self.device.vc_position = 3 + self.device.save() + + def _save_in_one_transaction(self, *saves): + """Run each of *saves* in order in one transaction; return the spy of the module reapplies.""" + with module_reapplies() as reapplies, self.captureOnCommitCallbacks(execute=True), transaction.atomic(): + for save in saves: + save() + return reapplies + + def _save_with_a_device_change(self, device_change, save, device_first, before=()): + """Run *device_change* and *save* in one transaction, *device_change* first when *device_first*. + + The saves in *before* run first in the same transaction. + """ + ordered = (device_change, save) if device_first else (save, device_change) + return self._save_in_one_transaction(*before, *ordered) + + +def flat_rule(module_type, name_template, **scope): + """Create a flat breakout rule with two channels for *module_type*.""" + return InterfaceNameRule.objects.create( + module_type=module_type, + name_template=name_template, + breakout_mode=BreakoutModeChoices.FLAT, + channel_count=2, + channel_start=0, + **scope, + ) + + +class BayEditTestCase(ModuleMoveTestCase): + """Install modules and edit their bays through real saves, with the committed callbacks run.""" + + @classmethod + def setUpTestData(cls): + super().setUpTestData() + cls.plain_type = cls._module_type("Plain", "{module}") + InterfaceNameRule.objects.create(module_type=cls.plain_type, name_template="et-{vc_position}/0/{bay_position}") + + @staticmethod + def _save_edit(bay, **values): + for field, value in values.items(): + setattr(bay, field, value) + bay.save() + + def _edit(self, bay, **values): + with self.captureOnCommitCallbacks(execute=True): + self._save_edit(bay, **values) diff --git a/netbox_interface_name_rules/transactions.py b/netbox_interface_name_rules/transactions.py index a70bbc27..f61007a3 100644 --- a/netbox_interface_name_rules/transactions.py +++ b/netbox_interface_name_rules/transactions.py @@ -1,49 +1,199 @@ # SPDX-License-Identifier: Apache-2.0 # Copyright (C) 2025 Marcin Zieba -"""Atomic blocks whose NetBox events count only when the block commits.""" +"""The one owner of the plugin's database connections and transaction state. -from contextlib import contextmanager +Every plugin write runs in a write scope. On main the scope holds ``default``. In a netbox-branching +branch it holds ``default`` and the branch alias: netbox-branching writes the ChangeDiff rows of a +branch change on ``default``, so a block opens both, ``default`` outside the branch. +""" + +from contextlib import ExitStack, contextmanager from contextvars import ContextVar +from dataclasses import dataclass +from functools import partial -from django.db import transaction +from django.db import DEFAULT_DB_ALIAS, DatabaseError, connections, router, transaction from netbox.context import events_queue +# PostgreSQL does not detect a lock cycle through the two sessions of one request in a branch. +LOCK_TIMEOUT = "10s" +READ_LOCK_TIMEOUT = "SHOW lock_timeout" +# Sets the session value, which the setup and each restore use. +SET_LOCK_TIMEOUT = "SELECT set_config('lock_timeout', %s, false)" + +# The aliases of the open write scope, or None outside one. +_scope_aliases = ContextVar("write_scope_aliases", default=None) # The event queues that enclose the current block, outermost first. _enclosing_queues = ContextVar("atomic_with_events_enclosing_queues", default=()) +def write_alias() -> str: + """Return the alias that NetBox's router gives for a write of an interface.""" + from dcim.models import Interface + + return router.db_for_write(Interface) + + @contextmanager -def atomic_with_events(using=None): - """Run the block in ``transaction.atomic()`` with its own NetBox event queue, kept only when the block commits. +def write_scope(expected_alias=None): + """Pin the aliases of one plugin operation and yield them: ``default``, then the branch alias in a branch. - A block that raises, or that ``transaction.set_rollback()`` rolls back, drops the events that it queued. + It raises before any query when the write alias is not *expected_alias*. A nested scope joins the + open one, and raises when its write alias differs. In a branch the outermost scope sets + ``lock_timeout`` on both connections and restores each value when it exits; on main it runs no query. """ - enclosing = (*_enclosing_queues.get(), events_queue.get()) - enclosing_token = _enclosing_queues.set(enclosing) - queue_token = events_queue.set({}) - outermost = transaction.get_autocommit(using=using) - commit_marks = [] - released = False + alias = write_alias() + if expected_alias is not None and alias != expected_alias: + raise RuntimeError(f"The write alias is {alias!r}, but the operation expects {expected_alias!r}.") + enclosing = _scope_aliases.get() + if enclosing is not None: + if alias != enclosing[-1]: + raise RuntimeError(f"The write alias is {alias!r}, but the open write scope writes to {enclosing[-1]!r}.") + yield enclosing + return + aliases = (DEFAULT_DB_ALIAS,) if alias == DEFAULT_DB_ALIAS else (DEFAULT_DB_ALIAS, alias) + token = _scope_aliases.set(aliases) try: - with transaction.atomic(using=using): - if outermost: - # The first commit callback: it runs after COMMIT, before a callback that can raise. - transaction.on_commit(lambda: commit_marks.append(True), using=using) - yield - rollback_marked = transaction.get_rollback(using=using) - released = not rollback_marked - finally: - events = events_queue.get() - events_queue.reset(queue_token) - _enclosing_queues.reset(enclosing_token) - # COMMIT runs only at the exit of the outermost block, and it can fail there or in a callback after it. - if commit_marks if outermost else released: - _keep(events) + if len(aliases) == 1: + yield aliases else: - _point_at_saved_rows(events, enclosing, using) + with _lock_timeout(aliases): + yield aliases + finally: + _scope_aliases.reset(token) + + +@contextmanager +def _lock_timeout(aliases): + """Set ``lock_timeout`` on the connection of each of *aliases*, and restore each value when the block exits. + + A caller's open transaction holds both the set and the restore, so either outcome of it keeps the + value from before it. A setup that fails restores the connections it changed. + """ + previous = {} + try: + for alias in aliases: + connection = connections[alias] + with connection.cursor() as cursor: + cursor.execute(READ_LOCK_TIMEOUT) + (value,) = cursor.fetchone() + cursor.execute(SET_LOCK_TIMEOUT, [LOCK_TIMEOUT]) + previous[alias] = value + yield + finally: + _restore_lock_timeouts(previous) + + +def _restore_lock_timeouts(previous): + """Restore each connection's ``lock_timeout`` from *previous*; discard each connection whose restore fails.""" + failures = [] + for alias, value in previous.items(): + connection = connections[alias] + try: + with connection.cursor() as cursor: + cursor.execute(SET_LOCK_TIMEOUT, [value]) + except DatabaseError as error: + _discard(connection) + failures.append((alias, error)) + if failures: + names = ", ".join(repr(alias) for alias, _ in failures) + raise RuntimeError(f"Could not restore lock_timeout on {names}; the connection is closed.") from failures[0][1] + + +def _discard(connection): + """Close the session of *connection*, so that no pool can hand out its changed setting.""" + session = connection.connection + try: + if session is not None: + session.close() + finally: + connection.close() + + +class _CommitJoin: + """Run a callback once, when every alias of the scope acknowledged the commit of its transaction.""" + + def __init__(self, callback, aliases): + self._callback = callback + # The whole set is pending before the first registration: a connection in autocommit acknowledges at once. + self._pending = set(aliases) + self._fired = False + + def acknowledge(self, alias): + self._pending.discard(alias) + if self._pending or self._fired: + return + self._fired = True + self._callback() + + +def on_commit(callback): + """Run *callback* once, after the transactions open at the call commit on every alias of the write scope. + + Each alias acknowledges through Django's ``on_commit``, which drops the acknowledgement when its + transaction, or a savepoint around the call, rolls back; *callback* then never runs. + """ + aliases = _scope_aliases.get() + if aliases is None: + raise RuntimeError("on_commit() runs only inside a write scope.") + if len(aliases) == 1: + transaction.on_commit(callback, using=aliases[0]) + return + join = _CommitJoin(callback, aliases) + for alias in aliases: + transaction.on_commit(partial(join.acknowledge, alias), using=alias) + + +@dataclass(frozen=True) +class Block: + """An open ``atomic_with_events()`` block, with one atomic block on each alias of its write scope.""" + + aliases: tuple[str, ...] + + def set_rollback(self): + """Mark the block for rollback on every alias.""" + for alias in self.aliases: + transaction.set_rollback(True, using=alias) + + +@contextmanager +def atomic_with_events(): + """Run the block atomically on every alias of the write scope, with its own NetBox event queue. + + It joins the open write scope or opens one. It opens ``default`` first and the branch alias inside + it, so the branch commits first. The block keeps its queued events when the branch block commits + (outermost) or is released (nested) with no alias marked for rollback, and drops them otherwise. + """ + with write_scope() as aliases: + enclosing = (*_enclosing_queues.get(), events_queue.get()) + enclosing_token = _enclosing_queues.set(enclosing) + queue_token = events_queue.set({}) + write = aliases[-1] + outermost = transaction.get_autocommit(using=write) + commit_marks = [] + released = False + try: + with ExitStack() as blocks: + for alias in aliases: + blocks.enter_context(transaction.atomic(using=alias)) + if outermost: + # The first commit callback of the write connection: it runs after COMMIT, before a callback can raise. + transaction.on_commit(lambda: commit_marks.append(True), using=write) + yield Block(aliases) + rollback_marked = any(transaction.get_rollback(using=alias) for alias in aliases) + released = not rollback_marked + finally: + events = events_queue.get() + events_queue.reset(queue_token) + _enclosing_queues.reset(enclosing_token) + # COMMIT runs only at the exit of the outermost block, and it can fail there or in a callback after it. + if commit_marks if outermost else released: + _keep(events) + else: + _point_at_saved_rows(events, enclosing) -def _point_at_saved_rows(dropped, enclosing, using): +def _point_at_saved_rows(dropped, enclosing): """Point each enclosing event of an object that the rolled-back block queued at a new copy of its saved row. NetBox 4.5 and later serialize the queued instance at the flush, and the block can have changed it. @@ -54,7 +204,7 @@ def _point_at_saved_rows(dropped, enclosing, using): if event["event_type"] == OBJECT_DELETED or "object" not in event: continue model = event["object_type"].model_class() - row = model._base_manager.using(using).filter(pk=event["object_id"]).first() + row = model._base_manager.filter(pk=event["object_id"]).first() if row is None: raise RuntimeError( f"{model._meta.label} {event['object_id']} has a queued event but no row after a rollback" diff --git a/netbox_interface_name_rules/views.py b/netbox_interface_name_rules/views.py index 7cb56421..7a4d467f 100644 --- a/netbox_interface_name_rules/views.py +++ b/netbox_interface_name_rules/views.py @@ -30,7 +30,7 @@ from .name_template import NamingContext, variables_for_context from .tables import InterfaceNameRuleTable from .template_variable_reference import naming_context_reference, rule_tester_variable_rows, variable_reference_rows -from .transactions import atomic_with_events +from .transactions import atomic_with_events, write_scope logger = logging.getLogger(__name__) @@ -41,6 +41,13 @@ APPLY_BATCH_LIMIT = 50 +def _failure_text(error): + """Return what the operator reads about *error*: each channel to rename back, or else the error type.""" + from .family import ChannelReconciliationError + + return str(error) if isinstance(error, ChannelReconciliationError) else type(error).__name__ + + @dataclasses.dataclass class RulePreview: """Lightweight stand-in for InterfaceNameRule used in the test/preview view.""" @@ -551,7 +558,8 @@ def get(self, request, **kwargs): # A conversion-scan failure must not blank the unrelated apply preview above. try: # Each family scanned costs a dry-run conversion, so the scan takes the same batch cap. - preview_conversions = find_convertible_families(rule, limit=APPLY_BATCH_LIMIT) + with write_scope(): + preview_conversions = find_convertible_families(rule, limit=APPLY_BATCH_LIMIT) except (re.error, ValueError) as exc: logger.exception("Failed to compute the conversion preview for rule %s", rule) messages.error(request, f"Failed to compute the conversion preview: {exc}") @@ -580,16 +588,13 @@ def get(self, request, **kwargs): ) def _enqueue(self, request, rule, job_class, name): - """Enqueue *job_class* against *rule* and report the outcome to the operator.""" + """Enqueue *job_class* against *rule* in the branch of *request*, and report the outcome to the operator.""" + from .jobs import rule_job_kwargs + + kwargs = rule_job_kwargs(rule.pk) try: - job = job_class.enqueue( - # instance is intentionally omitted: InterfaceNameRule does not - # inherit JobsMixin, so passing instance= would fail full_clean(). - # The job is still named and findable in Core → Jobs. - name=name, - user=request.user, - rule_id=rule.pk, - ) + # No instance=: the rule has no JobsMixin, so Job.full_clean() would refuse it. + job = job_class.enqueue(name=name, user=request.user, **kwargs) messages.success(request, f"Background job enqueued (job #{job.pk}). Check Core → Jobs for status.") except Exception as e: logger.exception("Failed to enqueue background job for rule %s", rule) @@ -613,10 +618,11 @@ def _convert(self, request, rule): ) return try: - outcome = convert_flat_families(rule, convert_ids) + with write_scope(): + outcome = convert_flat_families(rule, convert_ids) except Exception as e: logger.exception("Failed to convert families for rule %s", rule) - messages.error(request, f"Failed to convert families: {type(e).__name__}") + messages.error(request, f"Failed to convert families: {_failure_text(e)}") return converted = len(outcome.changed_families) messages.success(request, f"Converted {converted} interface family(ies) to the channelized topology.") @@ -649,7 +655,8 @@ def post(self, request, **kwargs): if not interface_ids: messages.warning(request, "No interfaces selected; nothing was applied.") else: - outcome = apply_rule_to_existing(rule, limit=APPLY_BATCH_LIMIT, interface_ids=interface_ids) + with write_scope(): + outcome = apply_rule_to_existing(rule, limit=APPLY_BATCH_LIMIT, interface_ids=interface_ids) messages.success(request, f"Applied rule: {outcome.changed_count} interface(s) renamed.") if outcome.skipped_members: messages.warning( @@ -658,7 +665,7 @@ def post(self, request, **kwargs): ) except Exception as e: logger.exception("Failed to apply rule %s", rule) - messages.error(request, f"Failed to apply rule {rule}: {type(e).__name__}") + messages.error(request, f"Failed to apply rule {rule}: {_failure_text(e)}") return redirect("plugins:netbox_interface_name_rules:interfacenamerule_apply_detail", pk=rule.pk) diff --git a/pyproject.toml b/pyproject.toml index bea6eb69..1ebc6408 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -103,7 +103,7 @@ devcontainer = [ [tool.pytest.ini_options] pythonpath = ["/opt/netbox/netbox", "."] DJANGO_SETTINGS_MODULE = "netbox_interface_name_rules.tests.isolated_settings" -addopts = "-n auto --dist loadscope --reuse-db --cov=netbox_interface_name_rules --cov-report=term-missing" +addopts = "-n auto --dist loadscope --reuse-db --cov=netbox_interface_name_rules --cov-report=term-missing --cov-fail-under=0" [tool.ruff] line-length = 120 @@ -211,6 +211,7 @@ branch = true [tool.coverage.report] fail_under = 97 +precision = 2 show_missing = true [tool.semantic_release]