diff --git a/docs/developer/utils.rst b/docs/developer/utils.rst index f2270fb7e..55aa8fb67 100644 --- a/docs/developer/utils.rst +++ b/docs/developer/utils.rst @@ -407,3 +407,93 @@ updates a ``WHOISInfo`` record. This signal is emitted when a WHOIS lookup is not triggered because the lookup conditions were not met (for example, an up-to-date WHOIS record already exists). + +.. _cache_invalidation: + +Cache Invalidation +------------------ + +OpenWISP Controller caches expensive values such as the configuration +checksum of devices and VPN servers, and the controller view responses. +When a *related* object changes (for example a certificate is renewed, a +device group context is edited, or a template is deleted) the cached value +can become stale and must be invalidated. + +This is handled by a declarative engine built around the +``CacheDependency`` class. Instead of scattering ``signal.connect()`` +calls across the codebase, each model that owns a cached value declares, +in one place, which related changes invalidate it. + +A model that owns a cache mixes in ``CacheInvalidationMixin`` and +overrides ``get_cache_dependencies()`` to return a list of +``CacheDependency`` objects: + +.. code-block:: python + + from openwisp_controller.config.base.cache import ( + CacheDependency, + CacheInvalidationMixin, + ) + + + class AbstractConfig(CacheInvalidationMixin, ...): + @classmethod + def get_cache_dependencies(cls): + return [ + # recompute the owning Config checksum when its client + # certificate changes + CacheDependency( + source="django_x509.Cert", + signal="post_save", + resolve=cls._resolve_cert_dependency, + target="update_status_if_checksum_changed", + ), + ] + +These declarations are wired at startup by +``register_cache_dependencies()``. Caches that are not owned by a model +(the controller view caches and the device group cache) are declared +instead in the ``AppConfig``, in +``ConfigConfig.connect_cache_dependencies()``. + +.. _print_cache_dependencies: + +``print_cache_dependencies`` +~~~~~~~~~~~~~~~~~~~~~~~~~~~~ + +Because dependencies are declared in more than one place, this management +command prints every cache dependency wired in the running project, so the +whole invalidation graph can be inspected at a glance: + +.. code-block:: bash + + python manage.py print_cache_dependencies + +The output is grouped by source and signal, and reports the target action, +the resolver, any tracked fields, and the dispatch UID of each dependency: + +.. code-block:: text + + config.device (post_save) + target: update_status_if_checksum_changed + resolve: _resolve_device_dependency track_fields: os, group_id, organization_id on_create: False on_commit: True + uid: cache_invalidation.config.config.config.device.post_save.update_status_if_checksum_changed._resolve_device_dependency.os+group_id+organization_id + + config.template (pre_delete) + target: update_status_if_checksum_changed + resolve: _resolve_template_dependency on_create: False on_commit: True + uid: cache_invalidation.config.config.config.template.pre_delete.update_status_if_checksum_changed._resolve_template_dependency + +Pass ``--format json`` for machine-readable output (useful, for example, +in a CI check that the wiring has not silently drifted): + +.. code-block:: bash + + python manage.py print_cache_dependencies --format json + +.. note:: + + This engine invalidates caches automatically when related objects + change. You still need to run ``clear_cache`` manually after editing + the :ref:`OPENWISP_CONTROLLER_CONTEXT ` setting, + because those system-wide variables are not tied to any model change. diff --git a/openwisp_controller/config/api/views.py b/openwisp_controller/config/api/views.py index e0ac2a9ca..5d6858b6a 100644 --- a/openwisp_controller/config/api/views.py +++ b/openwisp_controller/config/api/views.py @@ -40,7 +40,6 @@ Config = load_model("config", "Config") VpnClient = load_model("config", "VpnClient") Cert = load_model("django_x509", "Cert") -Organization = load_model("openwisp_users", "Organization") class TemplateListCreateView(ProtectedAPIMixin, ListCreateAPIView): @@ -274,14 +273,12 @@ def devicegroup_delete_invalidates_cache(cls, organization_id): cls._invalidate_from_queryset(qs) @classmethod - def certificate_delete_invalidates_cache(cls, organization_id, common_name): - try: - assert common_name - org_slug = Organization.objects.only("slug").get(id=organization_id).slug - except (AssertionError, Organization.DoesNotExist): + def certificate_delete_invalidates_cache(cls, common_name, organization_slug=None): + if not common_name: return cls.get_device_group.invalidate(cls, "", common_name) - cls.get_device_group.invalidate(cls, org_slug, common_name) + if organization_slug: + cls.get_device_group.invalidate(cls, organization_slug, common_name) template_list = TemplateListCreateView.as_view() diff --git a/openwisp_controller/config/apps.py b/openwisp_controller/config/apps.py index 4c7e233f4..032321db9 100644 --- a/openwisp_controller/config/apps.py +++ b/openwisp_controller/config/apps.py @@ -2,13 +2,7 @@ from django.conf import settings from django.core.exceptions import ImproperlyConfigured from django.db.models import Case, Count, When -from django.db.models.signals import ( - m2m_changed, - post_delete, - post_save, - pre_delete, - pre_save, -) +from django.db.models.signals import m2m_changed, post_delete, post_save, pre_save from django.urls import register_converter from django.utils.translation import gettext_lazy as _ from openwisp_notifications.types import ( @@ -53,12 +47,114 @@ def ready(self, *args, **kwargs): self.connect_signals() self.register_notification_types() self.add_ignore_notification_widget() - self.enable_cache_invalidation() + self.connect_related_changes_handlers() + self.connect_cache_dependencies() self.register_dashboard_charts() self.register_menu_groups() self.notification_cache_update() connect_whois_handlers() + def connect_cache_dependencies(self): + """ + Wires the declarative cache-invalidation dependencies. + + Models that own a cached value declare their related-change + dependencies in ``get_cache_dependencies`` (see + ``CacheInvalidationMixin``). + Caches that are not owned by a model (controller view caches and device + group caches) are declared here. Connecting all of them in one place + replaces the cache-invalidation ``signal.connect()`` calls that were + previously scattered across the codebase. + """ + from .base.cache import CacheDependency, _resolve_pk_snapshot + from .controller.views import DeviceChecksumView + from .handlers import ( + devicegroup_delete_handler, + invalidate_devicegroup_cache_change_handler, + organization_disabled_handler, + ) + + # Model-owned checksum caches (declared on the models themselves). + self.config_model.register_cache_dependencies() + self.vpn_model.register_cache_dependencies() + + dependencies = [ + # DeviceChecksumView caches are invalidated when a device is created, + # updated, deleted or when its config is deactivated. + CacheDependency( + source=self.device_model, + signal="post_save", + on_create=True, + target=DeviceChecksumView.invalidate_get_device_cache, + ), + # Deferred to commit so a concurrent request cannot repopulate the + # cache with a device that is about to be (or was just) deleted. + # ``post_delete`` + ``_resolve_pk_snapshot`` because Django clears + # ``instance.pk`` on deleted instances before the deferred + # on_commit callback runs (see ``_resolve_pk_snapshot``). + CacheDependency( + source=self.device_model, + signal="post_delete", + resolve=_resolve_pk_snapshot, + target=DeviceChecksumView.invalidate_get_device_cache, + ), + CacheDependency( + signal_obj=config_deactivated, + name="config_deactivated", + target=( + DeviceChecksumView.invalidate_get_device_cache_on_config_deactivated + ), + ), + # When an organization is disabled, all its devices are deactivated, + # so we need to invalidate the controller view caches for all objects. + CacheDependency( + source=self.org_model, + signal="pre_save", + on_commit=False, + target=organization_disabled_handler, + ), + # Invalidate the DeviceGroupCommonName cache when a device's group, + # a device group, or a certificate changes. + CacheDependency( + signal_obj=device_group_changed, + name="device_group_changed", + source=self.device_model, + target=invalidate_devicegroup_cache_change_handler, + ), + CacheDependency( + source=self.devicegroup_model, + signal="post_save", + target=invalidate_devicegroup_cache_change_handler, + ), + CacheDependency( + source=self.cert_model, + signal="post_save", + target=invalidate_devicegroup_cache_change_handler, + ), + # Kept synchronous (on_commit=False) so devicegroup_delete_handler + # still receives the live instance and can read organization_id + # before Django clears instance.pk post-delete. The handler + # itself defers the actual task enqueue via transaction.on_commit(). + CacheDependency( + source=self.devicegroup_model, + signal="post_delete", + on_commit=False, + target=devicegroup_delete_handler, + ), + # Same as above: kept synchronous so the handler can also read + # common_name before Django clears instance.pk post-delete. + CacheDependency( + source=self.cert_model, + signal="post_delete", + on_commit=False, + target=devicegroup_delete_handler, + ), + ] + for dependency in dependencies: + dependency.connect( + dispatch_uid=dependency.build_dispatch_uid("cache_invalidation.app") + ) + def __setmodels__(self): self.device_model = load_model("config", "Device") self.template_model = load_model("config", "Template") @@ -125,11 +221,6 @@ def connect_signals(self): sender=self.vpn_model, dispatch_uid="vpn.post_delete", ) - post_save.connect( - self.config_model.certificate_updated, - sender=self.cert_model, - dispatch_uid="cert_update_invalidate_checksum_cache", - ) group_templates_changed.connect( handlers.devicegroup_templates_change_handler, sender=self.devicegroup_model, @@ -150,11 +241,6 @@ def connect_signals(self): sender=self.template_model, dispatch_uid="template_pre_save_handler", ) - pre_save.connect( - handlers.organization_disabled_handler, - sender=self.org_model, - dispatch_uid="organization_disabled_pre_save_clear_device_checksum_cache", - ) post_save.connect( self.template_model.post_save_handler, sender=self.template_model, @@ -268,76 +354,32 @@ def add_ignore_notification_widget(self): obj_notification_widget, ) - def enable_cache_invalidation(self): + def connect_related_changes_handlers(self): """ - Triggers the cache invalidation for the - device config checksum (view and model method) + Connects signal handlers that react to a change in one object by + propagating side effects to related objects. These are intentionally + kept out of the declarative cache-invalidation engine (see + ``connect_cache_dependencies``) because they do more than invalidate a + cached value: + + * clearing a device's management IP when its config is deactivated; + * re-applying group templates when a device's group changes + (``devicegroup_change_handler``); + * refreshing the configs of a VPN server's clients when the server + changes. ``vpn_server_change_handler`` recomputes each client's + checksum and emits ``config_modified`` for it, but only when that + checksum actually changed. """ - from .controller.views import DeviceChecksumView, GetVpnView - from .handlers import ( - device_cache_invalidation_handler, - devicegroup_change_handler, - devicegroup_delete_handler, - vpn_server_change_handler, - ) + from .handlers import devicegroup_change_handler, vpn_server_change_handler - post_save.connect( - DeviceChecksumView.invalidate_get_device_cache, - sender=self.device_model, - dispatch_uid="invalidate_get_device_cache", - ) config_deactivated.connect( self.device_model.config_deactivated_clear_management_ip, dispatch_uid="config_deactivated_clear_management_ip", ) - config_deactivated.connect( - DeviceChecksumView.invalidate_get_device_cache_on_config_deactivated, - dispatch_uid="config_deactivated_invalidate_get_device_cache", - ) - # VPN cache invalidation - post_save.connect( - GetVpnView.invalidate_get_vpn_cache, - sender=self.vpn_model, - dispatch_uid="invalidate_get_vpn_cache", - ) - pre_delete.connect( - GetVpnView.invalidate_get_vpn_cache, - sender=self.vpn_model, - dispatch_uid="vpn_server_pre_delete_invalidate_get_vpn_cache", - ) - vpn_server_modified.connect( - GetVpnView.invalidate_get_vpn_cache, - dispatch_uid="vpn_server_modified_invalidate_get_vpn_cache", - ) device_group_changed.connect( devicegroup_change_handler, sender=self.device_model, - dispatch_uid="invalidate_devicegroup_cache_on_device_change", - ) - post_save.connect( - devicegroup_change_handler, - sender=self.devicegroup_model, - dispatch_uid="invalidate_devicegroup_cache_on_devicegroup_change", - ) - post_save.connect( - devicegroup_change_handler, - sender=self.cert_model, - dispatch_uid="invalidate_devicegroup_cache_on_certificate_change", - ) - post_delete.connect( - devicegroup_delete_handler, - sender=self.devicegroup_model, - dispatch_uid="invalidate_devicegroup_cache_on_devicegroup_delete", - ) - post_delete.connect( - devicegroup_delete_handler, - sender=self.cert_model, - dispatch_uid="invalidate_devicegroup_cache_on_certificate_delete", - ) - pre_delete.connect( - device_cache_invalidation_handler, - sender=self.device_model, - dispatch_uid="device.invalidate_cache", + dispatch_uid="manage_devicegroup_templates_on_device_change", ) vpn_server_modified.connect( vpn_server_change_handler, diff --git a/openwisp_controller/config/base/cache.py b/openwisp_controller/config/base/cache.py new file mode 100644 index 000000000..c8ebc4f36 --- /dev/null +++ b/openwisp_controller/config/base/cache.py @@ -0,0 +1,402 @@ +import json +from types import SimpleNamespace + +from django.core.exceptions import FieldDoesNotExist +from django.db import models, transaction +from django.db.models.signals import post_delete, post_save, pre_delete, pre_save +from django.utils.translation import gettext_lazy as _ +from swapper import load_model + +# Maps the string names used in declarations to the actual Django signals. +_MODEL_SIGNALS = { + "post_save": post_save, + "post_delete": post_delete, + "pre_delete": pre_delete, + "pre_save": pre_save, +} + + +def _default_resolve(instance, **kwargs): + """Default resolver: act on the instance that emitted the signal.""" + return [instance] + + +def _resolve_pk_snapshot(instance, **kwargs): + """ + Resolver for delete-triggered dependencies deferred via ``on_commit``. + + Django's ``Collector.delete()`` sets ``instance.pk`` to ``None`` on every + deleted instance immediately after ``pre_delete``/``post_delete`` signals + fire, well before an ``on_commit`` callback actually runs. Returning + ``[instance]`` here would hand the deferred callback a ``None`` pk. This + returns a disposable object exposing only the pk value, captured now + while it's still valid. + """ + return [SimpleNamespace(pk=instance.pk)] + + +class CacheDependency: + """ + Declarative description of a related change that must invalidate a cache. + + This is the single, generic mechanism used across the config app to keep + cached values (configuration checksums, controller view caches, device + group caches) in sync when a *related* object changes. + + A dependency is wired to a Django signal by :meth:`connect`. When the + signal fires, :attr:`resolve` returns the objects whose cache must be + invalidated and :attr:`target` is applied to each of them. ``target`` is + either the name of a method to call on each resolved object, or a callable + invoked with the resolved object. + + Parameters + ---------- + target: + Either a method name (``str``) invoked on each resolved object, or a + callable ``target(obj)``. Reusing the existing action methods (e.g. + ``update_status_if_checksum_changed``, ``invalidate_checksum_cache``) + and view classmethods keeps behavior identical. + resolve: + Callable ``resolve(instance, **signal_kwargs)`` returning an iterable + of the objects ``target`` must act on. Defaults to acting on the + instance that emitted the signal (``[instance]``). + source: + The signal sender. Either a swappable model label (e.g. + ``"django_x509.Cert"``) resolved lazily via ``swapper.load_model``, a + model class, or ``None`` (any sender). Ignored when ``signal_obj`` is a + custom signal that does not filter by sender. + signal: + One of ``post_save``, ``post_delete``, ``pre_delete``, ``pre_save``. + Ignored when ``signal_obj`` is provided. + signal_obj: + A custom Django ``Signal`` instance (e.g. ``config_deactivated``) to + connect to instead of one of the model signals above. + track_fields: + Optional iterable of source field names whose *value* must actually + change for the dependency to fire. Enabling this registers a + ``pre_save`` handler that snapshots the old values so the ``post_save`` + handler can compare them, mirroring the manual ``save()`` change + detection that some models used to perform. + on_create: + Whether to act when ``post_save`` reports ``created=True`` + (default ``False``). + on_commit: + Whether to defer ``target`` to ``transaction.on_commit`` + (default ``True``, matching the existing handlers). + """ + + _SNAPSHOT_ATTR = "_cache_dependency_snapshots" + + # Every connected dependency (model-owned or app-level) registers itself + # here, keyed by its dispatch_uid, so the whole invalidation graph can be + # introspected at runtime (see ``get_registered_dependencies`` and the + # ``print_cache_dependencies`` management command). + _registry = {} + + def __init__( + self, + *, + target, + resolve=_default_resolve, + source=None, + signal="post_save", + signal_obj=None, + name=None, + track_fields=None, + on_create=False, + on_commit=True, + ): + self.target = target + self.resolve = resolve + self.source = source + self.signal_name = signal + self.signal_obj = signal_obj + self.name = name + self.track_fields = list(track_fields) if track_fields else None + self.on_create = on_create + self.on_commit = on_commit + self._uid = None + + @property + def signal(self): + if self.signal_obj is not None: + return self.signal_obj + return _MODEL_SIGNALS[self.signal_name] + + @property + def sender(self): + if isinstance(self.source, str): + app_label, model_name = self.source.split(".") + return load_model(app_label, model_name) + return self.source + + def build_dispatch_uid(self, prefix): + """ + Builds a descriptive, order-independent ``dispatch_uid``. + + Deriving the uid from the sender, signal and target keeps it stable when + the surrounding dependency list is reordered and makes it readable in + tracebacks. ``name`` disambiguates custom signals, which have no natural + name of their own. The resolver, tracked fields and timing are also + encoded so two dependencies that share sender/signal/target but differ + in those attributes cannot collide and silently overwrite each other in + ``CacheDependency._registry``. + """ + sender = self.sender + sender_label = sender._meta.label_lower if sender is not None else "any" + if self.signal_obj is not None: + signal_label = self.name or "signal" + else: + signal_label = self.signal_name + target_label = ( + self.target if isinstance(self.target, str) else self.target.__name__ + ) + resolve_label = ( + "instance" if self.resolve is _default_resolve else self.resolve.__name__ + ) + parts = [prefix, sender_label, signal_label, target_label, resolve_label] + if self.track_fields: + parts.append("+".join(self.track_fields)) + if self.on_create: + parts.append("oncreate") + if not self.on_commit: + parts.append("immediate") + return ".".join(parts) + + def connect(self, dispatch_uid): + """Connect this dependency's handler to its signal.""" + self._uid = dispatch_uid + if self.track_fields: + pre_save.connect( + self._snapshot_handler, + sender=self.sender, + dispatch_uid=f"{dispatch_uid}.snapshot", + weak=False, + ) + self.signal.connect( + self._handler, + sender=self.sender, + dispatch_uid=dispatch_uid, + weak=False, + ) + CacheDependency._registry[dispatch_uid] = self + + def disconnect(self): + """Disconnect this dependency's handlers (useful for test isolation).""" + if self._uid is None: + return + if self.track_fields: + pre_save.disconnect( + sender=self.sender, dispatch_uid=f"{self._uid}.snapshot" + ) + self.signal.disconnect(sender=self.sender, dispatch_uid=self._uid) + CacheDependency._registry.pop(self._uid, None) + + @classmethod + def get_registered_dependencies(cls): + """Returns all connected dependencies, sorted by ``dispatch_uid``.""" + return [cls._registry[uid] for uid in sorted(cls._registry)] + + def describe(self): + """Returns a plain dict describing this dependency for introspection.""" + sender = self.sender + if self.signal_obj is not None: + signal = self.name or "custom" + else: + signal = self.signal_name + target = ( + self.target if isinstance(self.target, str) else self.target.__qualname__ + ) + if self.resolve is _default_resolve: + resolve = "instance" + else: + resolve = self.resolve.__name__ + return { + "source": sender._meta.label_lower if sender is not None else "any", + "signal": signal, + "target": target, + "resolve": resolve, + "track_fields": self.track_fields, + "on_create": self.on_create, + "on_commit": self.on_commit, + "dispatch_uid": self._uid, + } + + @classmethod + def render_registered(cls, fmt="text"): + """ + Returns a string describing every connected cache dependency. + + ``fmt`` is either ``"text"`` (human-readable, grouped by source and + signal) or ``"json"`` (a machine-readable list of ``describe()`` dicts). + """ + dependencies = cls.get_registered_dependencies() + if fmt == "json": + return json.dumps([dep.describe() for dep in dependencies], indent=2) + if not dependencies: + return _("No cache dependencies are registered.") + lines = [] + last_group = None + for dep in dependencies: + info = dep.describe() + group = (info["source"], info["signal"]) + if group != last_group: + if last_group is not None: + lines.append("") + lines.append("{0} ({1})".format(info["source"], info["signal"])) + last_group = group + lines.append(" " + _("target: {target}").format(target=info["target"])) + details = " " + _("resolve: {resolve}").format(resolve=info["resolve"]) + if info["track_fields"]: + details += _(" track_fields: {fields}").format( + fields=", ".join(info["track_fields"]) + ) + details += _(" on_create: {on_create} on_commit: {on_commit}").format( + on_create=info["on_create"], on_commit=info["on_commit"] + ) + lines.append(details) + lines.append(" " + _("uid: {uid}").format(uid=info["dispatch_uid"])) + return "\n".join(lines) + + def _snapshot_handler(self, sender, instance, **kwargs): + """Store the old values of ``track_fields`` before the instance saves.""" + if instance._state.adding or instance.pk is None: + return + if not self._may_track_fields_change(instance, **kwargs): + # A previous save may have left a snapshot on this instance; drop it + # so it cannot bleed into this save(update_fields=...) comparison. + self._discard_snapshot(instance) + return + snapshot = self._snapshot_from_initial_values(instance) + if snapshot is None: + snapshot = self._snapshot_from_db(sender, instance) + if snapshot is None: + self._discard_snapshot(instance) + return + snapshots = instance.__dict__.setdefault(self._SNAPSHOT_ATTR, {}) + snapshots[self._uid] = snapshot + + def _discard_snapshot(self, instance): + """Remove this dependency's stored snapshot from ``instance`` if any.""" + snapshots = getattr(instance, self._SNAPSHOT_ATTR, None) + if snapshots is not None: + snapshots.pop(self._uid, None) + + def _may_track_fields_change(self, instance, **kwargs): + """ + ``save(update_fields=[...])`` guarantees only those fields are + persisted; if none of them are ``track_fields``, nothing we care + about could have changed, so skip snapshotting (and the DB fetch + it may trigger) entirely. + """ + update_fields = kwargs.get("update_fields") + if update_fields is None: + return True + expanded = set(update_fields) + for name in update_fields: + try: + model_field = instance._meta.get_field(name) + except FieldDoesNotExist: + continue + expanded.add(model_field.name) + expanded.add(model_field.attname) + return any(field in expanded for field in self.track_fields) + + def _snapshot_from_initial_values(self, instance): + """ + Builds the snapshot from ``_initial_`` attributes already set + by the model (e.g. ``Device._set_initial_values_for_changed_checked_fields``), + avoiding a DB round-trip. Returns ``None`` if any tracked field lacks + one, so the caller falls back to fetching the old values from the DB. + """ + missing = object() + snapshot = dict() + for field in self.track_fields: + value = getattr(instance, f"_initial_{field}", missing) + if value is missing: + return None + snapshot[field] = value + return snapshot + + def _snapshot_from_db(self, sender, instance): + try: + old = sender._default_manager.only(*self.track_fields).get(pk=instance.pk) + except sender.DoesNotExist: + return None + return {field: getattr(old, field) for field in self.track_fields} + + def _tracked_fields_changed(self, instance): + snapshots = getattr(instance, self._SNAPSHOT_ATTR, None) or {} + # Consume the snapshot: a reused instance saved again (e.g. with + # ``update_fields``) must not compare against this stale snapshot. + old = snapshots.pop(self._uid, None) + if old is None: + # No snapshot (e.g. on creation) -> nothing to compare against. + return False + for field, old_value in old.items(): + if old_value is models.DEFERRED: + # Old value unknown (was deferred at snapshot time); assume changed. + return True + if old_value != getattr(instance, field): + return True + return False + + def _should_skip(self, instance, **kwargs): + if ( + self.signal is post_save + and kwargs.get("created", False) + and not self.on_create + ): + return True + if self.track_fields and not self._tracked_fields_changed(instance): + return True + return False + + def _apply(self, objects): + for obj in objects: + if obj is None: + continue + if callable(self.target): + self.target(obj) + else: + getattr(obj, self.target)() + + def _handler(self, sender, instance, **kwargs): + if self._should_skip(instance, **kwargs): + return + objects = self.resolve(instance, **kwargs) + if not objects: + return + objects = list(objects) + if self.on_commit: + transaction.on_commit(lambda: self._apply(objects)) + else: + self._apply(objects) + + +class CacheInvalidationMixin: + """ + Lets a cache-owning model declare, in one place, which related changes + invalidate its cached value(s). + + Subclasses override :meth:`get_cache_dependencies` to return a list of + :class:`CacheDependency`, and ``AppConfig.ready()`` calls + :meth:`register_cache_dependencies` to wire the Django signals. Adding a new + related-field dependency is then a matter of appending a declaration, + instead of scattering ``signal.connect()`` calls across the app. + + The declarations are returned by a classmethod (rather than held in a class + attribute) so they can reference the model's own private classmethods, which + do not exist yet while the class body is being evaluated. + """ + + @classmethod + def get_cache_dependencies(cls): + """Returns the list of :class:`CacheDependency` for this model.""" + return [] + + @classmethod + def register_cache_dependencies(cls): + prefix = f"cache_invalidation.{cls._meta.label_lower}" + for dependency in cls.get_cache_dependencies(): + dependency.connect(dispatch_uid=dependency.build_dispatch_uid(prefix)) diff --git a/openwisp_controller/config/base/config.py b/openwisp_controller/config/base/config.py index 7b2412b8f..10c2260aa 100644 --- a/openwisp_controller/config/base/config.py +++ b/openwisp_controller/config/base/config.py @@ -27,6 +27,7 @@ from ..sortedm2m.fields import SortedManyToManyField from ..utils import get_default_templates_queryset from .base import BaseConfig, ChecksumCacheMixin, get_cached_args_rewrite +from .cache import CacheDependency, CacheInvalidationMixin logger = logging.getLogger(__name__) @@ -40,7 +41,7 @@ def __str__(self): return _("Relationship with {0}").format(self.template.name) -class AbstractConfig(ChecksumCacheMixin, BaseConfig): +class AbstractConfig(CacheInvalidationMixin, ChecksumCacheMixin, BaseConfig): """ Abstract model implementing the NetJSON DeviceConfiguration object @@ -168,6 +169,128 @@ def get_cached_checksum(self): self.refresh_from_db(fields=["checksum_db"]) return self.checksum_db + @classmethod + def _resolve_cert_dependency(cls, cert, **kwargs): + """ + Returns the Config whose checksum depends on this client certificate. + + Mirrors the previous ``certificate_updated`` handler: a revoked + certificate or a certificate not linked to a VpnClient does not affect + any configuration. + """ + if cert.revoked: + return [] + try: + return [cert.vpnclient.config] + except ObjectDoesNotExist: + return [] + + @classmethod + def _bulk_invalidate_configs(cls, filters): + """Bulk-recompute the checksums of the configs matching ``filters``.""" + from ..tasks import bulk_invalidate_config_get_cached_checksum + + bulk_invalidate_config_get_cached_checksum.delay(filters) + + @classmethod + def _invalidate_configs_in_group(cls, group): + """Bulk-recompute checksums of configs in a group (context change).""" + cls._bulk_invalidate_configs({"device__group_id": str(group.id)}) + + @classmethod + def _invalidate_configs_in_org(cls, org_config_settings): + """Bulk-recompute checksums of an org's configs (context change).""" + cls._bulk_invalidate_configs( + {"device__organization_id": str(org_config_settings.organization_id)} + ) + + @classmethod + def _resolve_device_dependency(cls, device, **kwargs): + try: + return [device.config] + except ObjectDoesNotExist: + return [] + + @classmethod + def _resolve_template_dependency(cls, template, **kwargs): + """ + Return configs that use ``template`` (captured before cascade delete). + + Skipped when the delete originates from an Organization (e.g. + ``org.delete()``): a config using one of that organization's own + templates is necessarily in the same organization and will be + cascade-deleted in the same transaction, so there is nothing left + to invalidate. + + Note this check only excludes an Organization origin, it does not + require the origin to be the Template itself. Template also cascades + from ``Vpn`` (``vpn`` FK, ``on_delete=CASCADE``): deleting a VPN + removes its VPN-type templates with ``origin`` set to the ``Vpn`` + instance, not ``Template``. ``Config.templates`` is a many-to-many + field, so the configs using those templates are *not* deleted by + that cascade and still need their checksum recomputed. Narrowing + this to "only run when origin is Template" would silently skip that + case and leave those configs with a stale cached checksum. + """ + origin = kwargs.get("origin") + if origin is not None: + Organization = load_model("openwisp_users", "Organization") + origin_model = ( + origin.model if isinstance(origin, models.QuerySet) else type(origin) + ) + if issubclass(origin_model, Organization): + return [] + return list(cls.objects.filter(templates=template)) + + @classmethod + def get_cache_dependencies(cls): + return [ + # A client certificate's content (re-issue / key change) feeds into + # Config.get_vpn_context(); recompute the owning Config's checksum. + CacheDependency( + source="django_x509.Cert", + signal="post_save", + resolve=cls._resolve_cert_dependency, + target="update_status_if_checksum_changed", + ), + # Device.os feeds into Config._should_use_dsa(); Device.group_id + # and Device.organization_id determine group/org-level context. + # Recompute the owning Config's checksum when any of them changes. + CacheDependency( + source="config.Device", + signal="post_save", + track_fields=["os", "group_id", "organization_id"], + resolve=cls._resolve_device_dependency, + target="update_status_if_checksum_changed", + ), + # Group-level configuration variables feed into Config.get_context(); + # recompute checksums of all configs in the group when they change. + CacheDependency( + source="config.DeviceGroup", + signal="post_save", + track_fields=["context"], + target=cls._invalidate_configs_in_group, + ), + # Organization-level configuration variables feed into + # Config.get_context(); recompute checksums of all org configs. + CacheDependency( + source="config.OrganizationConfigSettings", + signal="post_save", + track_fields=["context"], + target=cls._invalidate_configs_in_org, + ), + # When a template is deleted, Django removes through-table + # rows without emitting m2m_changed. Capture the affected configs + # during pre_delete (while through rows still exist) and recompute + # their checksums on commit (after the cascade completes). + CacheDependency( + source="config.Template", + signal="pre_delete", + resolve=cls._resolve_template_dependency, + target="update_status_if_checksum_changed", + ), + ] + @classmethod def bulk_invalidate_get_cached_checksum(cls, query_params): """ @@ -464,17 +587,6 @@ def enforce_required_templates( *required_templates.order_by("name").values_list("pk", flat=True) ) - @classmethod - def certificate_updated(cls, instance, created, **kwargs): - if created or instance.revoked: - return - try: - config = instance.vpnclient.config - except ObjectDoesNotExist: - return - else: - transaction.on_commit(config.update_status_if_checksum_changed) - @classmethod def register_context_function(cls, func): """ diff --git a/openwisp_controller/config/base/device_group.py b/openwisp_controller/config/base/device_group.py index 7983c1fd9..5500e265a 100644 --- a/openwisp_controller/config/base/device_group.py +++ b/openwisp_controller/config/base/device_group.py @@ -15,7 +15,6 @@ from .. import settings as app_settings from ..signals import group_templates_changed from ..sortedm2m.fields import SortedManyToManyField -from ..tasks import bulk_invalidate_config_get_cached_checksum from .config import TemplatesThrough @@ -74,19 +73,6 @@ def clean(self): except SchemaError as e: raise ValidationError({"input": e.message}) - def save( - self, force_insert=False, force_update=False, using=None, update_fields=None - ): - context_changed = False - if not self._state.adding: - db_instance = self.__class__.objects.only("context").get(id=self.id) - context_changed = db_instance.context != self.context - super().save(force_insert, force_update, using, update_fields) - if context_changed: - bulk_invalidate_config_get_cached_checksum.delay( - {"device__group_id": str(self.id)} - ) - def get_context(self): return deepcopy(self.context) diff --git a/openwisp_controller/config/base/multitenancy.py b/openwisp_controller/config/base/multitenancy.py index 67b96761e..a1c7360f1 100644 --- a/openwisp_controller/config/base/multitenancy.py +++ b/openwisp_controller/config/base/multitenancy.py @@ -12,7 +12,6 @@ from .. import settings as app_settings from ..exceptions import OrganizationDeviceLimitExceeded -from ..tasks import bulk_invalidate_config_get_cached_checksum class AbstractOrganizationConfigSettings(UUIDModel): @@ -73,19 +72,6 @@ def clean(self): ) return super().clean() - def save( - self, force_insert=False, force_update=False, using=None, update_fields=None - ): - context_changed = False - if not self._state.adding: - db_instance = self.__class__.objects.only("context").get(id=self.id) - context_changed = db_instance.context != self.context - super().save(force_insert, force_update, using, update_fields) - if context_changed: - bulk_invalidate_config_get_cached_checksum.delay( - {"device__organization_id": str(self.organization_id)} - ) - class AbstractOrganizationLimits(models.Model): organization = models.OneToOneField( diff --git a/openwisp_controller/config/base/vpn.py b/openwisp_controller/config/base/vpn.py index f3270c852..61de40a66 100644 --- a/openwisp_controller/config/base/vpn.py +++ b/openwisp_controller/config/base/vpn.py @@ -15,6 +15,7 @@ from django.utils.functional import cached_property from django.utils.text import slugify from django.utils.translation import gettext_lazy as _ +from django_x509.signals import x509_renewed from swapper import get_model_name from openwisp_utils.base import KeyField @@ -34,6 +35,7 @@ trigger_zerotier_server_update_member, ) from .base import BaseConfig, ConfigChecksumCacheMixin +from .cache import CacheDependency, CacheInvalidationMixin, _resolve_pk_snapshot logger = logging.getLogger(__name__) @@ -43,7 +45,12 @@ def _peer_cache_key(vpn): return str(vpn.pk) -class AbstractVpn(ConfigChecksumCacheMixin, ShareableOrgMixinUniqueName, BaseConfig): +class AbstractVpn( + CacheInvalidationMixin, + ConfigChecksumCacheMixin, + ShareableOrgMixinUniqueName, + BaseConfig, +): """ Abstract VPN model """ @@ -307,15 +314,97 @@ def _check_changes(self): "private_key", "network_id", ] - current = self._meta.model.objects.only(*attrs).get(pk=self.pk) + current = self._meta.model.objects.only("name", *attrs).get(pk=self.pk) for attr in attrs: if getattr(self, attr) != getattr(current, attr): self._send_vpn_modified_after_save = True break + if not self._send_vpn_modified_after_save and self._is_backend_type("zerotier"): + if self.name != current.name: + self._send_vpn_modified_after_save = True def _send_vpn_modified_signal(self): vpn_server_modified.send(sender=self.__class__, instance=self) + def handle_related_change(self): + """ + Invalidates the VPN checksum and emits ``vpn_server_modified`` so that + client configs are recomputed. Called by :class:`CacheDependency` when + a related object (e.g. server CA/Cert) changes content. + """ + self.invalidate_checksum_cache() + self._send_vpn_modified_signal() + + @classmethod + def _invalidate_vpn_view_cache(cls, vpn): + """ + Invalidates the ``GetVpnView`` cache for a VPN server. Imported lazily + to avoid a circular import between models and ``controller.views``. + """ + from ..controller.views import GetVpnView + + GetVpnView.invalidate_get_vpn_cache(vpn) + + @classmethod + def _resolve_ca_dependency(cls, ca, **kwargs): + """Return VPNs whose server CA is ``ca``.""" + vpn_model = cls + return list(vpn_model.objects.filter(ca_id=ca.pk)) + + @classmethod + def _resolve_server_cert_dependency(cls, cert, **kwargs): + """Return VPNs whose server certificate is ``cert``.""" + vpn_model = cls + return list(vpn_model.objects.filter(cert_id=cert.pk)) + + @classmethod + def get_cache_dependencies(cls): + return [ + # The VPN server's own change (create or update) invalidates its + # controller view cache. + CacheDependency( + source="config.Vpn", + signal="post_save", + on_create=True, + target=cls._invalidate_vpn_view_cache, + ), + # Deferred to commit so a concurrent request cannot repopulate the + # cache with a VPN that is about to be (or was just) deleted. + # ``post_delete`` + ``_resolve_pk_snapshot`` because Django clears + # ``instance.pk`` on deleted instances before the deferred + # on_commit callback runs (see ``_resolve_pk_snapshot``). + CacheDependency( + source="config.Vpn", + signal="post_delete", + resolve=_resolve_pk_snapshot, + target=cls._invalidate_vpn_view_cache, + ), + # A change to the VPN server configuration (e.g. via related objects) + # emits ``vpn_server_modified`` and must invalidate the view cache. + CacheDependency( + signal_obj=vpn_server_modified, + name="vpn_server_modified", + target=cls._invalidate_vpn_view_cache, + ), + # When the server CA is renewed, the VPN's generated configuration + # changes; invalidate the VPN checksum and cascade to client configs. + CacheDependency( + source="django_x509.Ca", + signal_obj=x509_renewed, + name="x509_renewed", + resolve=cls._resolve_ca_dependency, + target="handle_related_change", + ), + # When the server certificate is renewed, same cascade as above. + CacheDependency( + source="django_x509.Cert", + signal_obj=x509_renewed, + name="x509_renewed", + resolve=cls._resolve_server_cert_dependency, + target="handle_related_change", + ), + ] + @classmethod def dhparam(cls, length): """ @@ -1053,10 +1142,25 @@ def _generate_zt_identity(self): @classmethod def invalidate_clients_cache(cls, vpn): """ - Invalidate checksum cache for clients that uses this VPN server + Recomputes the stored checksum of clients that use this VPN server. + + Changing a VPN server field (e.g. host, keys, subnet) alters the + context of every client configuration. Recompute each client's + checksum so that ``Config.checksum_db`` reflects the new VPN server + context, set its status to "modified" and emit ``config_modified`` + with action ``"related_template_changed"``. + + As with ``Config.update_status_if_checksum_changed()`` elsewhere, + ``config_modified`` is only emitted when the checksum actually + changed, not unconditionally for every client of the VPN server. """ - for client in vpn.vpnclient_set.iterator(): - # invalidate cache for device - client.config._send_config_modified_signal( - action="related_template_changed" - ) + for client in vpn.vpnclient_set.select_related( + "config", "config__device", "config__device__group" + ).iterator(): + config = client.config + # keep the historical signal action for this related change + config._config_modified_action = "related_template_changed" + if not config.update_status_if_checksum_changed(): + # no change: undo the action set above so it does not + # linger on this (otherwise throwaway) config instance + config._config_modified_action = "config_changed" diff --git a/openwisp_controller/config/handlers.py b/openwisp_controller/config/handlers.py index d838d1ad3..1443d0016 100644 --- a/openwisp_controller/config/handlers.py +++ b/openwisp_controller/config/handlers.py @@ -4,8 +4,6 @@ from openwisp_notifications.signals import notify from swapper import load_model -from openwisp_controller.config.controller.views import DeviceChecksumView - from . import tasks from .signals import config_status_changed, device_registered @@ -44,32 +42,83 @@ def device_registered_notification(sender, instance, is_new, **kwargs): def devicegroup_change_handler(instance, **kwargs): + """ + Manages group templates when a device's group changes. + + Cache invalidation for the device group change is handled separately + by ``invalidate_devicegroup_cache_change_handler``, declared as a + ``CacheDependency`` target in + ``ConfigConfig.connect_cache_dependencies`` (see ``config/apps.py``). + """ if type(instance) is list: # changes group templates for multiple devices devicegroup_templates_change_handler(instance, **kwargs) return if instance._state.adding or ("created" in kwargs and kwargs["created"] is True): return - model_name = instance._meta.model_name - if model_name == Device._meta.model_name: - # remove old group templates and apply new group templates - devicegroup_templates_change_handler(instance, **kwargs) - tasks.invalidate_devicegroup_cache_change.delay(instance.id, model_name) + # this handler is only connected to device_group_changed (sender=Device), + # so instance is always a Device here: remove old group templates and + # apply the new ones + devicegroup_templates_change_handler(instance, **kwargs) + + +def invalidate_devicegroup_cache_change_handler(instance, **kwargs): + """ + Invalidates the ``DeviceGroupCommonName`` cache when a device's group, + a device group, or a certificate changes. Used as a ``CacheDependency`` + target (see ``ConfigConfig.connect_cache_dependencies`` in + ``config/apps.py``). + """ + if isinstance(instance, list): + for device_id in instance: + tasks.invalidate_devicegroup_cache_change.delay( + device_id, Device._meta.model_name + ) + return + tasks.invalidate_devicegroup_cache_change.delay( + instance.id, instance._meta.model_name + ) def devicegroup_delete_handler(instance, **kwargs): + """ + Invalidates the ``DeviceGroupCommonName`` cache when a device group or a + certificate is deleted. Used as a ``CacheDependency`` target (see + ``ConfigConfig.connect_cache_dependencies`` in ``config/apps.py``). + + Runs synchronously (the ``CacheDependency`` is declared with + ``on_commit=False``) so it still receives the live ``instance``. Only the + task enqueue itself is deferred to ``transaction.on_commit()``, so a + concurrent request cannot repopulate the cache from a row that is about + to be (or was just) deleted. + + For a deleted ``Cert``, ``common_name`` and the organization's ``slug`` + are captured here rather than looked up by the deferred task: in an + organization-cascade delete, the ``Organization`` row itself is also + gone by the time the task would run, so it must not depend on a + post-commit database lookup to resolve the org's slug. A ``Cert`` can + have no organization at all (a cert shared across organizations), in + which case only the no-org cache entry is invalidated: ``get_device_group`` + only filters by organization when one is explicitly requested, so a + shared cert can still populate (and needs to invalidate) that entry. + """ kwargs = {} model_name = instance._meta.model_name - kwargs["organization_id"] = instance.organization_id if isinstance(instance, Cert): + if not instance.common_name: + return kwargs["common_name"] = instance.common_name - tasks.invalidate_devicegroup_cache_delete.delay(instance.id, model_name, **kwargs) - - -def device_cache_invalidation_handler(instance, **kwargs): - view = DeviceChecksumView() - setattr(view, "kwargs", {"pk": str(instance.pk)}) - view.get_device.invalidate(view) + organization = instance.organization + if organization is not None: + kwargs["organization_slug"] = organization.slug + else: + kwargs["organization_id"] = instance.organization_id + instance_id = instance.id + transaction.on_commit( + lambda: tasks.invalidate_devicegroup_cache_delete.delay( + instance_id, model_name, **kwargs + ) + ) def config_backend_change_handler(instance, **kwargs): diff --git a/openwisp_controller/config/management/commands/print_cache_dependencies.py b/openwisp_controller/config/management/commands/print_cache_dependencies.py new file mode 100644 index 000000000..febdf142b --- /dev/null +++ b/openwisp_controller/config/management/commands/print_cache_dependencies.py @@ -0,0 +1,22 @@ +from django.core.management.base import BaseCommand +from django.utils.translation import gettext_lazy as _ + +from openwisp_controller.config.base.cache import CacheDependency + + +class Command(BaseCommand): + help = _( + "Prints every cache dependency wired in the project, so the whole" + " cache invalidation graph can be inspected at a glance." + ) + + def add_arguments(self, parser): + parser.add_argument( + "--format", + choices=["text", "json"], + default="text", + help=_("Output format (default: text)."), + ) + + def handle(self, *args, **options): + self.stdout.write(CacheDependency.render_registered(fmt=options["format"])) diff --git a/openwisp_controller/config/tasks.py b/openwisp_controller/config/tasks.py index 798e40f8c..5ca2f7ad0 100644 --- a/openwisp_controller/config/tasks.py +++ b/openwisp_controller/config/tasks.py @@ -115,7 +115,7 @@ def invalidate_devicegroup_cache_delete(instance_id, model_name, **kwargs): ) elif model_name == Cert._meta.model_name: DeviceGroupCommonName.certificate_delete_invalidates_cache( - kwargs["organization_id"], kwargs["common_name"] + kwargs["common_name"], kwargs.get("organization_slug") ) diff --git a/openwisp_controller/config/tests/test_api.py b/openwisp_controller/config/tests/test_api.py index dbe9f94d5..e5f2136a1 100644 --- a/openwisp_controller/config/tests/test_api.py +++ b/openwisp_controller/config/tests/test_api.py @@ -3,6 +3,7 @@ from django.contrib.auth.models import Permission from django.core.cache import cache +from django.http.response import Http404 from django.test import TestCase from django.test.client import BOUNDARY, MULTIPART_CONTENT, encode_multipart from django.test.testcases import TransactionTestCase @@ -16,6 +17,7 @@ from openwisp_utils.tests import capture_any_output, catch_signal from .. import settings as app_settings +from ..controller.views import DeviceChecksumView, VpnChecksumView from ..signals import group_templates_changed from .utils import ( CreateConfigTemplateMixin, @@ -1487,6 +1489,32 @@ def test_devicegroup_commonname_cache_invalidates_on_cert_delete(self): response = self.client.get(path, data={"org": org.slug}) self.assertEqual(response.status_code, 404) + def test_device_and_vpn_view_cache_invalidate_on_organization_delete(self): + """ + Regression test for the Device and Vpn delete CacheDependency + declarations (``ConfigConfig.connect_cache_dependencies`` and + ``Vpn.get_cache_dependencies``): both are deferred to commit + (``post_delete`` + default ``on_commit=True``) and must keep working + when triggered as part of a larger Organization cascade delete, where + Django clears each deleted instance's pk before the on_commit + callback runs. + """ + _, org, _ = self._get_devicegroup_org_cert() + device = Device.objects.get(organization=org) + vpn = Vpn.objects.get(organization=org) + device_view = DeviceChecksumView() + device_view.kwargs = {"pk": str(device.pk)} + self.assertEqual(device_view.get_device(), device) + vpn_view = VpnChecksumView() + vpn_view.kwargs = {"pk": str(vpn.pk)} + self.assertEqual(vpn_view.get_vpn(), vpn) + + org.delete() + with self.assertRaises(Http404): + device_view.get_device() + with self.assertRaises(Http404): + vpn_view.get_vpn() + def test_devicegroup_templates_change(self): org = self._get_org() t1 = self._create_template(name="t1", organization=org) diff --git a/openwisp_controller/config/tests/test_config.py b/openwisp_controller/config/tests/test_config.py index 0d3f6772e..be157d984 100644 --- a/openwisp_controller/config/tests/test_config.py +++ b/openwisp_controller/config/tests/test_config.py @@ -1,7 +1,12 @@ +import json +import uuid from copy import deepcopy -from unittest.mock import patch +from io import StringIO +from unittest.mock import Mock, call, patch from django.core.exceptions import ValidationError +from django.core.management import call_command +from django.db import models from django.db.transaction import atomic from django.test import TestCase from django.test.testcases import TransactionTestCase @@ -11,7 +16,10 @@ from openwisp_utils.tests import catch_signal from .. import settings as app_settings +from .. import tasks from ..base.base import logger as base_config_logger +from ..base.cache import CacheDependency +from ..handlers import invalidate_devicegroup_cache_change_handler from ..signals import config_backend_changed, config_modified, config_status_changed from .utils import ( CreateConfigTemplateMixin, @@ -22,9 +30,12 @@ Config = load_model("config", "Config") Device = load_model("config", "Device") +DeviceGroup = load_model("config", "DeviceGroup") +OrganizationConfigSettings = load_model("config", "OrganizationConfigSettings") Template = load_model("config", "Template") Vpn = load_model("config", "Vpn") Ca = load_model("django_x509", "Ca") +Cert = load_model("django_x509", "Cert") class TestConfig( @@ -573,12 +584,15 @@ def test_certificate_updated_skipped_for_deactivated_config(self): self.assertEqual(config.status, "deactivating") # VpnClient is deleted on deactivation; cert is auto-revoked. self.assertEqual(config.vpnclient_set.count(), 0) - # Un-revoke the cert so certificate_updated() bypasses the early + # Un-revoke the cert so _resolve_cert_dependency() bypasses the early # "if revoked: return" guard and hits the ObjectDoesNotExist path. + # The Cert CacheDependency defers to transaction.on_commit, which + # TestCase does not fire unless captured. cert.revoked = False - cert.save() - # Config status must not change: certificate_updated() returns early - # because the VpnClient was deleted during deactivation. + with self.captureOnCommitCallbacks(execute=True): + cert.save() + # Config status must not change: _resolve_cert_dependency() returns + # early because the VpnClient was deleted during deactivation. config.refresh_from_db() self.assertEqual(config.status, "deactivating") @@ -925,6 +939,138 @@ def test_config_backend_changed(self): self.assertTrue(d.config.templates.filter(pk=t2.pk).exists()) self.assertFalse(d.config.templates.filter(pk=t1.pk).exists()) + def test_devicegroup_context_change_defers_checksum_invalidation_to_commit(self): + """ + Regression test for the DeviceGroup CacheDependency (deferred to + commit, default on_commit=True): the Celery task recomputing the + group's configs must be enqueued only after the transaction that + changed the group's context has committed, otherwise a worker + picking up the task can read the DB before the commit and wrongly + conclude the context did not change. + """ + device_group = self._create_device_group(context={"interface_type": "eth0"}) + with patch( + "openwisp_controller.config.tasks" + ".bulk_invalidate_config_get_cached_checksum.delay" + ) as mocked_delay: + with self.captureOnCommitCallbacks(execute=True): + device_group.context = {"interface_type": "eth1"} + device_group.full_clean() + device_group.save() + mocked_delay.assert_not_called() + mocked_delay.assert_called_once_with( + {"device__group_id": str(device_group.id)} + ) + + def test_organization_context_change_defers_checksum_invalidation_to_commit(self): + """ + Same as the DeviceGroup regression test above, but for the + OrganizationConfigSettings CacheDependency. + """ + org_settings = OrganizationConfigSettings.objects.create( + organization=self._get_org(), context={"interface_type": "eth0"} + ) + with patch( + "openwisp_controller.config.tasks" + ".bulk_invalidate_config_get_cached_checksum.delay" + ) as mocked_delay: + with self.captureOnCommitCallbacks(execute=True): + org_settings.context = {"interface_type": "eth1"} + org_settings.full_clean() + org_settings.save() + mocked_delay.assert_not_called() + mocked_delay.assert_called_once_with( + {"device__organization_id": str(org_settings.organization_id)} + ) + + def test_devicegroup_delete_invalidates_cache_deferred_to_commit(self): + """ + Regression test for the DeviceGroup post_delete CacheDependency + (config/apps.py), which targets ``devicegroup_delete_handler``: the + Celery task invalidating the DeviceGroupCommonName cache must be + enqueued only after the deleting transaction has committed, + otherwise a concurrent request can repopulate the cache from a + device group that is about to be deleted, leaving it stale after + commit. + """ + org = self._get_org() + device_group = self._create_device_group(organization=org) + device_group_id = device_group.id + with patch( + "openwisp_controller.config.tasks" + ".invalidate_devicegroup_cache_delete.delay" + ) as mocked_delay: + with self.captureOnCommitCallbacks(execute=True): + device_group.delete() + mocked_delay.assert_not_called() + mocked_delay.assert_called_once_with( + device_group_id, + DeviceGroup._meta.model_name, + organization_id=org.id, + ) + + def test_cert_delete_invalidates_devicegroup_cache_deferred_to_commit(self): + """ + Same as above, but for the Cert post_delete CacheDependency, which + also targets ``devicegroup_delete_handler``. + """ + org = self._get_org() + cert = self._create_cert(organization=org) + cert_id = cert.id + common_name = cert.common_name + with patch( + "openwisp_controller.config.tasks" + ".invalidate_devicegroup_cache_delete.delay" + ) as mocked_delay: + with self.captureOnCommitCallbacks(execute=True): + cert.delete() + mocked_delay.assert_not_called() + mocked_delay.assert_called_once_with( + cert_id, + Cert._meta.model_name, + common_name=common_name, + organization_slug=org.slug, + ) + + def test_shared_cert_delete_invalidates_devicegroup_wildcard_cache(self): + """ + A Cert with organization=None ("shared" cert, usable across + multiple organizations) is still reachable via the no-org (``""``) + DeviceGroupCommonName cache key: ``get_device_group`` only filters + by organization when one is explicitly requested. Deleting the cert + must still invalidate that wildcard entry, even though there is no + organization_slug to also invalidate an org-scoped entry. + """ + cert = self._create_cert(organization=None) + cert_id = cert.id + common_name = cert.common_name + with patch( + "openwisp_controller.config.tasks" + ".invalidate_devicegroup_cache_delete.delay" + ) as mocked_delay: + with self.captureOnCommitCallbacks(execute=True): + cert.delete() + mocked_delay.assert_not_called() + mocked_delay.assert_called_once_with( + cert_id, + Cert._meta.model_name, + common_name=common_name, + ) + + def test_shared_cert_delete_task_invalidates_devicegroup_wildcard_cache(self): + cert = self._create_cert(organization=None) + common_name = cert.common_name + with patch( + "openwisp_controller.config.api.views.DeviceGroupCommonName" + ".certificate_delete_invalidates_cache" + ) as mocked_invalidate: + tasks.invalidate_devicegroup_cache_delete( + cert.id, + Cert._meta.model_name, + common_name=common_name, + ) + mocked_invalidate.assert_called_once_with(common_name, None) + class TestTransactionConfig( CreateConfigTemplateMixin, @@ -1008,6 +1154,96 @@ def test_certificate_renew_invalidates_checksum_cache(self): config.refresh_from_db() self.assertEqual(config.status, "modified") + def test_device_os_change_updates_config_checksum(self): + org = self._get_org() + device = self._create_device( + name="test", organization=org, os="OpenWrt 19.07.0" + ) + config = self._create_config( + device=device, + backend="netjsonconfig.OpenWrt", + config={ + "interfaces": [ + { + "name": "eth0", + "type": "ethernet", + "addresses": [{"proto": "dhcp", "family": "ipv4"}], + } + ] + }, + ) + config.set_status_applied() + config.refresh_from_db() + old_checksum_db = config.checksum_db + self.assertEqual(config.status, "applied") + # changing the OS toggles DSA (disabled on 19.x, enabled on 21.x), + # which changes the rendered configuration + device.os = "OpenWrt 21.02.0" + device.save() + config = Config.objects.get(pk=config.pk) + self.assertEqual(config.status, "modified") + self.assertNotEqual(config.checksum_db, old_checksum_db) + self.assertEqual(config.checksum_db, config.checksum) + + def test_device_org_change_updates_config_checksum(self): + org1 = self._get_org() + OrganizationConfigSettings.objects.create( + organization=org1, context={"interface_type": "ethernet"} + ) + org2 = self._create_org(name="org2", slug="org2") + OrganizationConfigSettings.objects.create( + organization=org2, context={"interface_type": "virtual"} + ) + device = self._create_device(name="test", organization=org1) + template = self._create_template( + config={"interfaces": [{"name": "eth0", "type": "{{ interface_type }}"}]}, + default_values={"interface_type": "ethernet"}, + ) + config = self._create_config(device=device) + config.templates.add(template) + config.set_status_applied() + config.refresh_from_db() + old_checksum_db = config.checksum_db + self.assertEqual(config.status, "applied") + # changing the organization changes the org-level context, + # which changes the rendered configuration + device.organization = org2 + device.save() + config = Config.objects.get(pk=config.pk) + self.assertEqual(config.status, "modified") + self.assertNotEqual(config.checksum_db, old_checksum_db) + self.assertEqual(config.checksum_db, config.checksum) + + def test_device_group_change_updates_config_checksum(self): + org = self._get_org() + group1 = DeviceGroup( + name="group1", organization=org, context={"interface_type": "ethernet"} + ) + group1.full_clean() + group1.save() + group2 = DeviceGroup( + name="group2", organization=org, context={"interface_type": "virtual"} + ) + group2.full_clean() + group2.save() + device = self._create_device(name="test", organization=org, group=group1) + template = self._create_template( + config={"interfaces": [{"name": "eth0", "type": "{{ interface_type }}"}]}, + default_values={"interface_type": "ethernet"}, + ) + config = self._create_config(device=device) + config.templates.add(template) + config.set_status_applied() + config.refresh_from_db() + old_checksum_db = config.checksum_db + self.assertEqual(config.status, "applied") + device.group = group2 + device.save() + config = Config.objects.get(pk=config.pk) + self.assertEqual(config.status, "modified") + self.assertNotEqual(config.checksum_db, old_checksum_db) + self.assertEqual(config.checksum_db, config.checksum) + def test_checksum_db_accounts_for_vpnclient(self): vpn = self._create_wireguard_vpn() vpn_template = self._create_template( @@ -1018,3 +1254,366 @@ def test_checksum_db_accounts_for_vpnclient(self): config.refresh_from_db() config._invalidate_backend_instance_cache() self.assertEqual(config.checksum, config.checksum_db) + + def test_deleting_template_invalidates_config_checksum(self): + template = self._create_template( + name="test-template", + config={"interfaces": [{"name": "eth0", "type": "ethernet"}]}, + ) + config = self._create_config(device=self._create_device()) + config.templates.add(template) + config.set_status_applied() + config.refresh_from_db() + old_checksum_db = config.checksum_db + self.assertEqual(config.status, "applied") + template.delete() + config = Config.objects.get(pk=config.pk) + self.assertNotEqual(config.checksum_db, old_checksum_db) + self.assertEqual(config.checksum_db, config.checksum) + self.assertEqual(config.status, "modified") + + def test_bulk_deleting_templates_invalidates_config_checksum(self): + template1 = self._create_template( + name="test-template1", + config={"interfaces": [{"name": "eth0", "type": "ethernet"}]}, + ) + template2 = self._create_template( + name="test-template2", + config={"interfaces": [{"name": "eth1", "type": "ethernet"}]}, + ) + config1 = self._create_config(device=self._create_device(name="device1")) + config1.templates.add(template1) + config2 = self._create_config( + device=self._create_device(name="device2", mac_address="00:11:22:33:44:66") + ) + config2.templates.add(template2) + for config in (config1, config2): + config.set_status_applied() + config1.refresh_from_db() + config2.refresh_from_db() + old_checksum_db1 = config1.checksum_db + old_checksum_db2 = config2.checksum_db + self.assertEqual(config1.status, "applied") + self.assertEqual(config2.status, "applied") + # bulk delete via a queryset, not per-instance .delete() calls + Template.objects.filter(pk__in=[template1.pk, template2.pk]).delete() + config1 = Config.objects.get(pk=config1.pk) + config2 = Config.objects.get(pk=config2.pk) + self.assertNotEqual(config1.checksum_db, old_checksum_db1) + self.assertEqual(config1.checksum_db, config1.checksum) + self.assertEqual(config1.status, "modified") + self.assertNotEqual(config2.checksum_db, old_checksum_db2) + self.assertEqual(config2.checksum_db, config2.checksum) + self.assertEqual(config2.status, "modified") + + def test_deleting_vpn_invalidates_config_checksum(self): + device, vpn, _ = self._create_wireguard_vpn_template() + config = device.config + config.set_status_applied() + config.refresh_from_db() + old_checksum_db = config.checksum_db + self.assertEqual(config.status, "applied") + # deleting the VPN cascades to delete its VPN-type template + # (Template.vpn is on_delete=CASCADE); the config using that + # template is not deleted (Config.templates is a many-to-many + # field), so its checksum must still be recomputed. + vpn.delete() + self.assertEqual(Template.objects.count(), 0) + config = Config.objects.get(pk=config.pk) + self.assertNotEqual(config.checksum_db, old_checksum_db) + self.assertEqual(config.checksum_db, config.checksum) + self.assertEqual(config.status, "modified") + + +class TestCacheDependency(CreateConfigTemplateMixin, CreateDeviceGroupMixin, TestCase): + """ + Unit tests for the declarative cache-invalidation engine + (``CacheDependency``) that centralizes cache/checksum invalidation + (issue #1095). + """ + + def _connect(self, **kwargs): + dependency = CacheDependency(**kwargs) + dependency.connect(dispatch_uid="test.cache_dependency") + self.addCleanup(dependency.disconnect) + return dependency + + def test_target_invoked_on_related_change(self): + target = Mock() + self._connect( + source="config.DeviceGroup", + signal="post_save", + on_commit=False, + resolve=lambda instance, **kwargs: [instance], + target=target, + ) + # creation is skipped by default (on_create=False) + group = self._create_device_group() + target.assert_not_called() + # an update fires the dependency with the resolved object + group.name = "renamed" + group.save() + target.assert_called_once_with(group) + + def test_on_create_opt_in(self): + target = Mock() + self._connect( + source="config.DeviceGroup", + signal="post_save", + on_create=True, + on_commit=False, + resolve=lambda instance, **kwargs: [instance], + target=target, + ) + group = self._create_device_group() + target.assert_called_once_with(group) + + def test_track_fields_fires_only_on_value_change(self): + target = Mock() + self._connect( + source="config.DeviceGroup", + signal="post_save", + track_fields=["context"], + on_commit=False, + resolve=lambda instance, **kwargs: [instance], + target=target, + ) + group = self._create_device_group(context={"a": "1"}) + target.assert_not_called() + with self.subTest("save without changing tracked field"): + group.name = "renamed" + group.save() + target.assert_not_called() + with self.subTest("save changing tracked field"): + group.context = {"a": "2"} + group.save() + target.assert_called_once_with(group) + + def test_snapshot_not_reused_across_saves(self): + """ + A snapshot from a prior save must be consumed so a later + save(update_fields=...) that excludes the tracked field cannot compare + against the stale snapshot and fire the target a second time. + """ + target = Mock() + self._connect( + source="config.DeviceGroup", + signal="post_save", + track_fields=["context"], + on_commit=False, + resolve=lambda instance, **kwargs: [instance], + target=target, + ) + group = self._create_device_group(context={"a": "1"}) + group.context = {"a": "2"} + group.save() + target.assert_called_once_with(group) + # a subsequent save that does not touch the tracked field must not + # re-fire the target because of a leftover snapshot + target.reset_mock() + group.name = "renamed" + group.save(update_fields=["name"]) + target.assert_not_called() + + def test_target_as_method_name(self): + # a string target is invoked as a method on each resolved object + dependency = CacheDependency( + source="config.DeviceGroup", + signal="post_save", + resolve=lambda instance, **kwargs: [instance], + target="some_method", + ) + obj = Mock() + dependency._apply([obj]) + obj.some_method.assert_called_once_with() + + def test_snapshot_uses_initial_values_when_all_tracked_fields_available(self): + dependency = CacheDependency( + source="config.Device", + signal="post_save", + track_fields=["name", "organization_id"], + target=Mock(), + ) + dependency._uid = "test.cache_dependency.snapshot.initial" + device = self._create_device(name="device-initial") + old_name = device._initial_name + old_org_id = device._initial_organization_id + # Emulate pre_save state where current values may differ from initial ones. + device.name = "device-updated" + + with patch.object( + dependency, + "_snapshot_from_initial_values", + wraps=dependency._snapshot_from_initial_values, + ) as initial_spy, patch.object( + dependency, + "_snapshot_from_db", + wraps=dependency._snapshot_from_db, + ) as db_spy: + dependency._snapshot_handler(Device, device) + + initial_spy.assert_called_once_with(device) + db_spy.assert_not_called() + snapshot = device._cache_dependency_snapshots[dependency._uid] + self.assertEqual(snapshot["name"], old_name) + self.assertEqual(snapshot["organization_id"], old_org_id) + + def test_snapshot_falls_back_to_db_when_initial_fields_are_missing(self): + dependency = CacheDependency( + source="config.DeviceGroup", + signal="post_save", + track_fields=["context"], + target=Mock(), + ) + dependency._uid = "test.cache_dependency.snapshot.db_fallback" + group = self._create_device_group(context={"a": "1"}) + + with patch.object( + dependency, + "_snapshot_from_initial_values", + wraps=dependency._snapshot_from_initial_values, + ) as initial_spy, patch.object( + dependency, + "_snapshot_from_db", + wraps=dependency._snapshot_from_db, + ) as db_spy: + dependency._snapshot_handler(DeviceGroup, group) + + initial_spy.assert_called_once_with(group) + db_spy.assert_called_once_with(DeviceGroup, group) + snapshot = group._cache_dependency_snapshots[dependency._uid] + self.assertEqual(snapshot["context"], {"a": "1"}) + + def test_snapshot_falls_back_to_db_when_initial_fields_are_deferred(self): + dependency = CacheDependency( + source="config.Device", + signal="post_save", + track_fields=["name", "organization_id"], + target=Mock(), + ) + dependency._uid = "test.cache_dependency.snapshot.deferred" + device = self._create_device(name="device-deferred") + deferred_device = Device.objects.only("id").get(pk=device.pk) + self.assertEqual(deferred_device._initial_name, models.DEFERRED) + self.assertEqual(deferred_device._initial_organization_id, models.DEFERRED) + + with patch.object( + dependency, + "_snapshot_from_initial_values", + wraps=dependency._snapshot_from_initial_values, + ) as initial_spy, patch.object( + dependency, + "_snapshot_from_db", + wraps=dependency._snapshot_from_db, + ) as db_spy: + dependency._snapshot_handler(Device, deferred_device) + + initial_spy.assert_called_once_with(deferred_device) + db_spy.assert_not_called() + snapshot = deferred_device._cache_dependency_snapshots[dependency._uid] + self.assertEqual(snapshot["name"], models.DEFERRED) + self.assertEqual(snapshot["organization_id"], models.DEFERRED) + + def test_snapshot_skips_when_update_fields_excludes_tracked_fields(self): + dependency = CacheDependency( + source="config.Device", + signal="post_save", + track_fields=["os", "group_id", "organization_id"], + target=Mock(), + ) + dependency._uid = "test.cache_dependency.snapshot.skip_irrelevant_update_fields" + device = self._create_device(os="OpenWrt 22.03") + + with patch.object( + dependency, + "_snapshot_from_initial_values", + wraps=dependency._snapshot_from_initial_values, + ) as initial_spy, patch.object( + dependency, + "_snapshot_from_db", + wraps=dependency._snapshot_from_db, + ) as db_spy: + dependency._snapshot_handler( + Device, device, update_fields={"management_ip", "last_ip"} + ) + + initial_spy.assert_not_called() + db_spy.assert_not_called() + snapshots = getattr(device, dependency._SNAPSHOT_ATTR, {}) + self.assertNotIn(dependency._uid, snapshots) + + def test_registry_tracks_connect_and_disconnect(self): + uid = "test.cache_dependency.registry" + dependency = CacheDependency( + source="config.DeviceGroup", signal="post_save", target=Mock() + ) + self.assertNotIn(uid, CacheDependency._registry) + dependency.connect(dispatch_uid=uid) + self.assertIs(CacheDependency._registry[uid], dependency) + dependency.disconnect() + self.assertNotIn(uid, CacheDependency._registry) + + def test_registry_includes_core_dependencies(self): + # the app-level and model-owned dependencies are wired at app startup + targets = { + dep.describe()["target"] + for dep in CacheDependency.get_registered_dependencies() + } + self.assertIn("update_status_if_checksum_changed", targets) + self.assertIn("DeviceChecksumView.invalidate_get_device_cache", targets) + self.assertIn("AbstractVpn._invalidate_vpn_view_cache", targets) + + def test_describe_reports_dependency_attributes(self): + dependency = self._connect( + source="config.Device", + signal="post_save", + track_fields=["os", "group_id", "organization_id"], + resolve=Config._resolve_device_dependency, + target="update_status_if_checksum_changed", + ) + info = dependency.describe() + self.assertEqual(info["source"], Device._meta.label_lower) + self.assertEqual(info["signal"], "post_save") + self.assertEqual(info["target"], "update_status_if_checksum_changed") + self.assertEqual(info["resolve"], "_resolve_device_dependency") + self.assertEqual(info["track_fields"], ["os", "group_id", "organization_id"]) + self.assertEqual(info["on_commit"], True) + self.assertEqual(info["dispatch_uid"], "test.cache_dependency") + + def test_print_cache_dependencies_text(self): + out = StringIO() + call_command("print_cache_dependencies", stdout=out) + output = out.getvalue() + self.assertIn("update_status_if_checksum_changed", output) + self.assertIn("on_commit", output) + self.assertIn("uid:", output) + + def test_print_cache_dependencies_json(self): + out = StringIO() + call_command("print_cache_dependencies", format="json", stdout=out) + data = json.loads(out.getvalue()) + self.assertGreater(len(data), 0) + expected_keys = { + "source", + "signal", + "target", + "resolve", + "track_fields", + "on_create", + "on_commit", + "dispatch_uid", + } + for item in data: + self.assertEqual(set(item.keys()), expected_keys) + + def test_invalidate_devicegroup_cache_change_handler_bulk_list(self): + device_id1, device_id2 = uuid.uuid4(), uuid.uuid4() + with patch.object(tasks.invalidate_devicegroup_cache_change, "delay") as delay: + invalidate_devicegroup_cache_change_handler([device_id1, device_id2]) + self.assertEqual(delay.call_count, 2) + delay.assert_has_calls( + [ + call(device_id1, Device._meta.model_name), + call(device_id2, Device._meta.model_name), + ] + ) diff --git a/openwisp_controller/config/tests/test_controller.py b/openwisp_controller/config/tests/test_controller.py index 225512029..2fb85fe2a 100644 --- a/openwisp_controller/config/tests/test_controller.py +++ b/openwisp_controller/config/tests/test_controller.py @@ -5,7 +5,7 @@ from django.core.exceptions import ValidationError from django.db import connection from django.http.response import Http404 -from django.test import TestCase +from django.test import TestCase, TransactionTestCase from django.test.utils import CaptureQueriesContext from django.urls import reverse, reverse_lazy from swapper import load_model @@ -89,35 +89,6 @@ def _test_view_organization_disabled( response = method(url, {"key": obj.key}) self.assertEqual(response.status_code, 404) - def _test_deactivating_deactivated_device_view( - self, url_name, method="get", data=None - ): - data = data or {} - device = self._create_device_config() - config = device.config - # The endpoint returns 200 when config.status is modified - config.set_status_modified() - path = reverse(f"controller:{url_name}", args=[device.pk]) - payload = {"key": device.key, **data} - response = getattr(self.client, method)(path, payload) - self.assertEqual(response.status_code, 200) - - # The endpoint returns 200 when config.status is deactivating - config.set_status_deactivating() - path = reverse("controller:device_checksum", args=[device.pk]) - response = self.client.get(path, {"key": device.key}) - self.assertEqual(response.status_code, 200) - config.refresh_from_db() - self.assertEqual(config.status, "deactivating") - - # The endpoint returns 404 when config.status is deactivated - config.set_status_deactivated() - path = reverse("controller:device_checksum", args=[device.pk]) - response = self.client.get(path, {"key": device.key}) - self.assertEqual(response.status_code, 404) - config.refresh_from_db() - self.assertEqual(config.status, "deactivated") - def test_device_checksum(self): d = self._create_device_config() c = d.config @@ -138,128 +109,6 @@ def test_device_checksum(self): self.assertIsNotNone(d.last_ip) self.assertIsNone(d.management_ip) - def test_device_get_object_cached(self): - d = self._create_device_config() - view = DeviceChecksumView() - view.kwargs = {"pk": str(d.pk)} - logger = controller_views_logger - - with self.subTest("check cache set"): - with patch("django.core.cache.cache.set") as mock: - self.assertEqual(view.get_device(), d) - mock.assert_called_once() - - with self.subTest("check cache get"): - with patch("django.core.cache.cache.get", return_value=d) as mock: - view.get_device() - mock.assert_called_once() - - with self.subTest("ensure DB is hit when cache is clear"): - with patch.object(logger, "debug") as mocked_debug: - with self.assertNumQueries(1): - self.assertEqual(view.get_device(), d) - mocked_debug.assert_called_once() - - with self.subTest("ensure DB is NOT hit when cache is present"): - with patch.object(logger, "debug") as mocked_debug: - with self.assertNumQueries(0): - self.assertEqual(view.get_device(), d) - mocked_debug.assert_not_called() - - with self.subTest("test manual invalidation"): - with patch.object(logger, "debug") as mocked_debug: - with self.assertNumQueries(1): - view.get_device.invalidate(view) - self.assertEqual(view.get_device(), d) - mocked_debug.assert_called_once() - - with self.subTest("test automatic cache invalidation"): - with patch.object(logger, "debug") as mocked_debug: - d.os = "test_cache" - d.save() - mocked_debug.assert_called_once() - - with self.assertNumQueries(1): - obj = view.get_device() - self.assertEqual(obj.os, "test_cache") - - with self.subTest("test cache invalidation on device delete"): - d.delete(check_deactivated=False) - with self.assertNumQueries(1): - with self.assertRaises(Http404): - view.get_device() - - def test_get_cached_checksum(self): - d = self._create_device_config() - # avoid cache to be invalidated by the update of the addresses - d.last_ip = "127.0.0.1" - d.management_ip = "10.0.0.2" - d.save() - - url = reverse("controller:device_checksum", args=[d.pk]) - - with self.subTest("first request does not return value from cache"): - with self.assertNumQueries(3): - with patch.object( - controller_views_logger, "debug" - ) as mocked_view_debug: - with patch.object( - base_config_logger, "debug" - ) as mocked_config_debug: - self.client.get( - url, {"key": d.key, "management_ip": "10.0.0.2"} - ) - self.assertEqual(mocked_config_debug.call_count, 0) - self.assertEqual(mocked_view_debug.call_count, 1) - - with self.subTest("update_last_ip updates the cache"): - with self.assertNumQueries(3): - with patch.object( - controller_views_logger, "debug" - ) as mocked_view_debug: - with patch.object( - base_config_logger, "debug" - ) as mocked_config_debug: - self.client.get( - url, {"key": d.key, "management_ip": "10.0.0.3"} - ) - mocked_config_debug.assert_not_called() - mocked_view_debug.assert_called_once_with( - f"invalidated view cache for device ID {d.pk.hex}" - ) - view = DeviceChecksumView() - view.kwargs = {"pk": str(d.pk)} - key = view.get_device.get_cache_key(view) - d.refresh_from_db() - cached_device = cache.get(key) - self.assertEqual(cached_device, d) - self.assertEqual(cached_device.management_ip, "10.0.0.3") - - with self.subTest("second request returns value from cache"): - with self.assertNumQueries(0): - with patch.object( - controller_views_logger, "debug" - ) as mocked_view_debug: - with patch.object( - base_config_logger, "debug" - ) as mocked_config_debug: - self.client.get( - url, {"key": d.key, "management_ip": "10.0.0.3"} - ) - mocked_config_debug.assert_not_called() - mocked_view_debug.assert_not_called() - - with self.subTest("ensure cache invalidation works"): - old_checksum = d.config.checksum - with patch.object(base_config_logger, "debug") as mocked_debug: - d.config.config["general"]["timezone"] = "Europe/Rome" - d.config.full_clean() - d.config.save() - del d.config.backend_instance - self.assertNotEqual(d.config.checksum, old_checksum) - self.assertEqual(d.config.get_cached_checksum(), d.config.checksum) - mocked_debug.assert_called_once() - def test_device_checksum_requested_signal_is_emitted(self): d = self._create_device_config() url = reverse("controller:device_checksum", args=[d.pk]) @@ -291,9 +140,6 @@ def test_device_config_download_requested_signal_is_emitted(self): request=response.wsgi_request, ) - def test_device_config_deactivated_checksum(self): - self._test_deactivating_deactivated_device_view("device_checksum") - @capture_any_output() def test_device_checksum_400(self): d = self._create_device_config() @@ -334,9 +180,6 @@ def test_device_download_config(self): self.assertIsNotNone(d.last_ip) self.assertIsNone(d.management_ip) - def test_deactivated_device_download_config(self): - self._test_deactivating_deactivated_device_view("device_download_config") - def test_device_download_config_bad_uuid(self): d = self._create_device_config() valid = reverse("controller:device_download_config", args=[d.pk]) @@ -452,65 +295,18 @@ def test_vpn_get_object_cached(self): view.get_vpn.invalidate(view) mock.assert_called_once() - def test_vpn_checksum_cache_invalidation_handler(self): + def test_vpn_cache_invalidation_on_delete(self): vpn = self._create_vpn() - url = reverse("controller:vpn_checksum", args=[vpn.pk]) - # Warm up the cache - with self.assertNumQueries(1): - response = self.client.get(url, {"key": vpn.key}) - self.assertEqual(response.content.decode(), vpn.checksum) - - # Cache works are expected - with self.assertNumQueries(0): - response = self.client.get(url, {"key": vpn.key}) - self.assertEqual(response.content.decode(), vpn.checksum) - - # Change VPN config to trigger invalidation - vpn.config["openvpn"][0]["proto"] = "tcp-server" - vpn.full_clean() - vpn.save() - - del vpn.backend_instance - with self.assertNumQueries(1): - response = self.client.get(url, {"key": vpn.key}) - self.assertEqual(response.content.decode(), vpn.checksum) - - def test_vpn_download_config(self): - v = self._create_vpn() - url = reverse("controller:vpn_download_config", args=[v.pk]) - # First request will populate the cache - with self.assertNumQueries(1), patch.object( - Vpn, "generate", return_value=v.generate() - ) as mocked_generate: - response = self.client.get(url, {"key": v.key}) - mocked_generate.assert_called_once() - self.assertEqual( - response["Content-Disposition"], "attachment; filename=test.tar.gz" - ) - self._check_header(response) - - with self.subTest("Second request will return cached config"): - with patch.object(Vpn, "generate") as mocked_generate: - response = self.client.get(url, {"key": v.key}) - mocked_generate.assert_not_called() - self.assertEqual( - response["Content-Disposition"], "attachment; filename=test.tar.gz" - ) - self._check_header(response) - - with self.subTest("Changing Vpn configuration will invalidate cache"): - v.config["wireguard"][0]["port"] = "51821" - v.full_clean() - v.save() - with self.assertNumQueries(1), patch.object( - Vpn, "generate", return_value=v.generate() - ) as mocked_generate: - response = self.client.get(url, {"key": v.key}) - mocked_generate.assert_called_once() - self.assertEqual( - response["Content-Disposition"], "attachment; filename=test.tar.gz" - ) - self._check_header(response) + view = VpnChecksumView() + view.kwargs = {"pk": str(vpn.pk)} + # warm up the view cache + self.assertEqual(view.get_vpn(), vpn) + key = view.get_vpn.get_cache_key(view) + self.assertEqual(cache.get(key), vpn) + # deleting the VPN must invalidate the cached view object + with self.captureOnCommitCallbacks(execute=True): + vpn.delete() + self.assertEqual(cache.get(key), None) def test_vpn_download_config_bad_uuid(self): v = self._create_vpn() @@ -980,11 +776,6 @@ def test_device_report_status_error(self): self.assertEqual(d.config.status, "error") self.assertEqual(d.config.error_reason, error_reason) - def test_deactivated_device_report_status(self): - self._test_deactivating_deactivated_device_view( - "device_report_status", method="post", data={"status": "applied"} - ) - def test_device_report_status_bad_uuid(self): d = self._create_device_config() valid_url = reverse("controller:device_report_status", args=[d.pk]) @@ -1503,72 +1294,21 @@ def test_register_403_disabled_registration_setting(self): ).count() self.assertEqual(count, 0) - @patch.object(app_settings, "SHARED_MANAGEMENT_IP_ADDRESS_SPACE", False) - def test_ip_fields_not_duplicated(self): - org1 = self._get_org() - c1 = self._create_config(organization=org1) - d2 = self._create_device( - organization=org1, name="testdup", mac_address="00:11:22:33:66:77" - ) - c2 = self._create_config(device=d2) - org2 = self._create_org(name="org2", shared_secret="123456") - c3 = self._create_config(organization=org2) - with self.assertNumQueries(6): - self.client.get( - reverse("controller:device_checksum", args=[c3.device.pk]), - {"key": c3.device.key, "management_ip": "192.168.1.99"}, - ) - with self.assertNumQueries(6): - self.client.get( - reverse("controller:device_checksum", args=[c1.device.pk]), - {"key": c1.device.key, "management_ip": "192.168.1.99"}, - ) - with self.assertNumQueries(0): - # repeat the request to test the checksum view cache interaction - self.client.get( - reverse("controller:device_checksum", args=[c1.device.pk]), - {"key": c1.device.key, "management_ip": "192.168.1.99"}, - ) - # triggers more queries because devices with conflicting addresses - # need to be updated, luckily it does not happen often - with self.assertNumQueries(8): - self.client.get( - reverse("controller:device_checksum", args=[c2.device.pk]), - {"key": c2.device.key, "management_ip": "192.168.1.99"}, - ) - c1.refresh_from_db() - c2.refresh_from_db() - c3.refresh_from_db() - # device previously having the IP now won't have it anymore - self.assertNotEqual(c1.device.last_ip, c2.device.last_ip) - self.assertNotEqual(c1.device.management_ip, c2.device.management_ip) - self.assertIsNone(c1.device.management_ip) - self.assertEqual(c2.device.management_ip, "192.168.1.99") - # other organization is not affected - self.assertEqual(c3.device.last_ip, "127.0.0.1") - self.assertEqual(c3.device.management_ip, "192.168.1.99") - - with self.subTest("test interaction with DeviceChecksumView caching"): - view = DeviceChecksumView() - view.kwargs = {"pk": str(c1.device.pk)} - cached_device1 = view.get_device() - self.assertIsNone(cached_device1.management_ip) - - @patch.object(app_settings, "WHOIS_CONFIGURED", True) - def test_remove_duplicated_last_ip_no_nplus1_queries(self): - # When WHOIS is configured, clearing a duplicate's last_ip runs - # process_ip_data_and_location during save(), which reads - # device.organization via is_whois_enabled. If organization is not - # selected upfront, each duplicate triggers extra SELECTs (N+1). - # The marginal cost per duplicate must be a single UPDATE only. - def _measure(n): - ip = f"192.168.40.{n}" - org = self._create_org(name=f"dupes-{n}", shared_secret=f"dupes-secret-{n}") - incoming = self._create_device( - organization=org, - name=f"incoming-{n}", - mac_address=f"00:11:22:33:{n:02x}:99", - last_ip=ip, + @patch.object(app_settings, "WHOIS_CONFIGURED", True) + def test_remove_duplicated_last_ip_no_nplus1_queries(self): + # When WHOIS is configured, clearing a duplicate's last_ip runs + # process_ip_data_and_location during save(), which reads + # device.organization via is_whois_enabled. If organization is not + # selected upfront, each duplicate triggers extra SELECTs (N+1). + # The marginal cost per duplicate must be a single UPDATE only. + def _measure(n): + ip = f"192.168.40.{n}" + org = self._create_org(name=f"dupes-{n}", shared_secret=f"dupes-secret-{n}") + incoming = self._create_device( + organization=org, + name=f"incoming-{n}", + mac_address=f"00:11:22:33:{n:02x}:99", + last_ip=ip, ) for i in range(n): self._create_device( @@ -1663,3 +1403,295 @@ def test_device_registered_signal(self): handler.assert_called_once_with( sender=Device, signal=device_registered, instance=device, is_new=True ) + + +class TestControllerTransaction( + TestRegistrationMixin, + CreateConfigTemplateMixin, + TestVpnX509Mixin, + TransactionTestCase, +): + """ + TransactionTestCase variant for tests that exercise cache + invalidation handlers deferred to ``transaction.on_commit``. + """ + + def _check_header(self, response): + self.assertEqual(response["X-Openwisp-Controller"], "true") + + def _test_deactivating_deactivated_device_view( + self, url_name, method="get", data=None + ): + data = data or {} + device = self._create_device_config() + config = device.config + # The endpoint returns 200 when config.status is modified + config.set_status_modified() + path = reverse(f"controller:{url_name}", args=[device.pk]) + payload = {"key": device.key, **data} + response = getattr(self.client, method)(path, payload) + self.assertEqual(response.status_code, 200) + + # The endpoint returns 200 when config.status is deactivating + config.set_status_deactivating() + path = reverse("controller:device_checksum", args=[device.pk]) + response = self.client.get(path, {"key": device.key}) + self.assertEqual(response.status_code, 200) + config.refresh_from_db() + self.assertEqual(config.status, "deactivating") + + # The endpoint returns 404 when config.status is deactivated + config.set_status_deactivated() + path = reverse("controller:device_checksum", args=[device.pk]) + response = self.client.get(path, {"key": device.key}) + self.assertEqual(response.status_code, 404) + config.refresh_from_db() + self.assertEqual(config.status, "deactivated") + + def test_device_config_deactivated_checksum(self): + self._test_deactivating_deactivated_device_view("device_checksum") + + def test_deactivated_device_download_config(self): + self._test_deactivating_deactivated_device_view("device_download_config") + + def test_deactivated_device_report_status(self): + self._test_deactivating_deactivated_device_view( + "device_report_status", method="post", data={"status": "applied"} + ) + + def test_device_get_object_cached(self): + d = self._create_device_config() + view = DeviceChecksumView() + view.kwargs = {"pk": str(d.pk)} + logger = controller_views_logger + + with self.subTest("check cache set"): + with patch("django.core.cache.cache.set") as mock: + self.assertEqual(view.get_device(), d) + mock.assert_called_once() + + with self.subTest("check cache get"): + with patch("django.core.cache.cache.get", return_value=d) as mock: + view.get_device() + mock.assert_called_once() + + with self.subTest("ensure DB is hit when cache is clear"): + with patch.object(logger, "debug") as mocked_debug: + with self.assertNumQueries(1): + self.assertEqual(view.get_device(), d) + mocked_debug.assert_called_once() + + with self.subTest("ensure DB is NOT hit when cache is present"): + with patch.object(logger, "debug") as mocked_debug: + with self.assertNumQueries(0): + self.assertEqual(view.get_device(), d) + mocked_debug.assert_not_called() + + with self.subTest("test manual invalidation"): + with patch.object(logger, "debug") as mocked_debug: + with self.assertNumQueries(1): + view.get_device.invalidate(view) + self.assertEqual(view.get_device(), d) + mocked_debug.assert_called_once() + + with self.subTest("test automatic cache invalidation"): + with patch.object(logger, "debug") as mocked_debug: + d.os = "test_cache" + d.save() + mocked_debug.assert_called_once() + + with self.assertNumQueries(1): + obj = view.get_device() + self.assertEqual(obj.os, "test_cache") + + with self.subTest("test cache invalidation on device delete"): + d.delete(check_deactivated=False) + with self.assertNumQueries(1): + with self.assertRaises(Http404): + view.get_device() + + def test_get_cached_checksum(self): + d = self._create_device_config() + # avoid cache to be invalidated by the update of the addresses + d.last_ip = "127.0.0.1" + d.management_ip = "10.0.0.2" + d.save() + + url = reverse("controller:device_checksum", args=[d.pk]) + + with self.subTest("first request does not return value from cache"): + with self.assertNumQueries(3): + with patch.object( + controller_views_logger, "debug" + ) as mocked_view_debug: + with patch.object( + base_config_logger, "debug" + ) as mocked_config_debug: + self.client.get( + url, {"key": d.key, "management_ip": "10.0.0.2"} + ) + self.assertEqual(mocked_config_debug.call_count, 0) + self.assertEqual(mocked_view_debug.call_count, 1) + + with self.subTest("update_last_ip updates the cache"): + with self.assertNumQueries(3): + with patch.object( + controller_views_logger, "debug" + ) as mocked_view_debug: + with patch.object( + base_config_logger, "debug" + ) as mocked_config_debug: + self.client.get( + url, {"key": d.key, "management_ip": "10.0.0.3"} + ) + mocked_config_debug.assert_not_called() + mocked_view_debug.assert_called_once_with( + f"invalidated view cache for device ID {d.pk.hex}" + ) + view = DeviceChecksumView() + view.kwargs = {"pk": str(d.pk)} + key = view.get_device.get_cache_key(view) + d.refresh_from_db() + cached_device = cache.get(key) + self.assertEqual(cached_device, d) + self.assertEqual(cached_device.management_ip, "10.0.0.3") + + with self.subTest("second request returns value from cache"): + with self.assertNumQueries(0): + with patch.object( + controller_views_logger, "debug" + ) as mocked_view_debug: + with patch.object( + base_config_logger, "debug" + ) as mocked_config_debug: + self.client.get( + url, {"key": d.key, "management_ip": "10.0.0.3"} + ) + mocked_config_debug.assert_not_called() + mocked_view_debug.assert_not_called() + + with self.subTest("ensure cache invalidation works"): + old_checksum = d.config.checksum + with patch.object(base_config_logger, "debug") as mocked_debug: + d.config.config["general"]["timezone"] = "Europe/Rome" + d.config.full_clean() + d.config.save() + del d.config.backend_instance + self.assertNotEqual(d.config.checksum, old_checksum) + self.assertEqual(d.config.get_cached_checksum(), d.config.checksum) + mocked_debug.assert_called_once() + + @patch.object(app_settings, "SHARED_MANAGEMENT_IP_ADDRESS_SPACE", False) + def test_ip_fields_not_duplicated(self): + org1 = self._get_org() + c1 = self._create_config(organization=org1) + d2 = self._create_device( + organization=org1, name="testdup", mac_address="00:11:22:33:66:77" + ) + c2 = self._create_config(device=d2) + org2 = self._create_org(name="org2", shared_secret="123456") + c3 = self._create_config(organization=org2) + with self.assertNumQueries(6): + self.client.get( + reverse("controller:device_checksum", args=[c3.device.pk]), + {"key": c3.device.key, "management_ip": "192.168.1.99"}, + ) + with self.assertNumQueries(6): + self.client.get( + reverse("controller:device_checksum", args=[c1.device.pk]), + {"key": c1.device.key, "management_ip": "192.168.1.99"}, + ) + with self.assertNumQueries(0): + # repeat the request to test the checksum view cache interaction + self.client.get( + reverse("controller:device_checksum", args=[c1.device.pk]), + {"key": c1.device.key, "management_ip": "192.168.1.99"}, + ) + # triggers more queries because devices with conflicting addresses + # need to be updated, luckily it does not happen often + with self.assertNumQueries(8): + self.client.get( + reverse("controller:device_checksum", args=[c2.device.pk]), + {"key": c2.device.key, "management_ip": "192.168.1.99"}, + ) + c1.refresh_from_db() + c2.refresh_from_db() + c3.refresh_from_db() + # device previously having the IP now won't have it anymore + self.assertNotEqual(c1.device.last_ip, c2.device.last_ip) + self.assertNotEqual(c1.device.management_ip, c2.device.management_ip) + self.assertIsNone(c1.device.management_ip) + self.assertEqual(c2.device.management_ip, "192.168.1.99") + # other organization is not affected + self.assertEqual(c3.device.last_ip, "127.0.0.1") + self.assertEqual(c3.device.management_ip, "192.168.1.99") + + with self.subTest("test interaction with DeviceChecksumView caching"): + view = DeviceChecksumView() + view.kwargs = {"pk": str(c1.device.pk)} + cached_device1 = view.get_device() + self.assertIsNone(cached_device1.management_ip) + + def test_vpn_checksum_cache_invalidation_handler(self): + vpn = self._create_vpn() + url = reverse("controller:vpn_checksum", args=[vpn.pk]) + # Warm up the cache + with self.assertNumQueries(1): + response = self.client.get(url, {"key": vpn.key}) + self.assertEqual(response.content.decode(), vpn.checksum) + + # Cache works are expected + with self.assertNumQueries(0): + response = self.client.get(url, {"key": vpn.key}) + self.assertEqual(response.content.decode(), vpn.checksum) + + # Change VPN config to trigger invalidation + vpn.config["openvpn"][0]["proto"] = "tcp-server" + vpn.full_clean() + vpn.save() + + del vpn.backend_instance + with self.assertNumQueries(1): + response = self.client.get(url, {"key": vpn.key}) + self.assertEqual(response.content.decode(), vpn.checksum) + + def test_vpn_download_config(self): + v = self._create_vpn() + url = reverse("controller:vpn_download_config", args=[v.pk]) + # First request will populate the cache + with ( + self.assertNumQueries(1), + patch.object(Vpn, "generate", return_value=v.generate()) as mocked_generate, + ): + response = self.client.get(url, {"key": v.key}) + mocked_generate.assert_called_once() + self.assertEqual( + response["Content-Disposition"], "attachment; filename=test.tar.gz" + ) + self._check_header(response) + + with self.subTest("Second request will return cached config"): + with patch.object(Vpn, "generate") as mocked_generate: + response = self.client.get(url, {"key": v.key}) + mocked_generate.assert_not_called() + self.assertEqual( + response["Content-Disposition"], "attachment; filename=test.tar.gz" + ) + self._check_header(response) + + with self.subTest("Changing Vpn configuration will invalidate cache"): + v.config["wireguard"][0]["port"] = "51821" + v.full_clean() + v.save() + with ( + self.assertNumQueries(1), + patch.object( + Vpn, "generate", return_value=v.generate() + ) as mocked_generate, + ): + response = self.client.get(url, {"key": v.key}) + mocked_generate.assert_called_once() + self.assertEqual( + response["Content-Disposition"], "attachment; filename=test.tar.gz" + ) + self._check_header(response) diff --git a/openwisp_controller/config/tests/test_device.py b/openwisp_controller/config/tests/test_device.py index 1e2bd4a3e..ba15c0103 100644 --- a/openwisp_controller/config/tests/test_device.py +++ b/openwisp_controller/config/tests/test_device.py @@ -382,57 +382,6 @@ def test_configuration_variable_priority(self, *args): config.get_context()["variable_type"], "global-context-variable" ) - def test_changing_org_variable_invalidates_cache(self): - org = self._get_org() - config_settings = OrganizationConfigSettings.objects.create( - organization=org, context={} - ) - config = self._create_config(organization=org) - template = self._create_template( - config={"interfaces": [{"name": "eth0", "type": "{{ interface_type }}"}]}, - default_values={"interface_type": "ethernet"}, - ) - config.templates.add(template) - old_checksum = config.get_cached_checksum() - - # Changing OrganizationConfigSettings.context should invalidate the cache - # and a new checksum should be returned in the following request. - config_settings.context = {"interface_type": "virtual"} - config_settings.full_clean() - config_settings.save() - - # Config.backend_instance is a "cached_property", hence deleting - # the attribute is required for testing. - del config.backend_instance - - new_checksum = config.get_cached_checksum() - self.assertNotEqual(old_checksum, new_checksum) - - def test_changing_group_variable_invalidates_cache(self): - org = self._get_org() - device_group = self._create_device_group(organization=org, context={}) - device = self._create_device(organization=org, group=device_group) - config = self._create_config(device=device) - template = self._create_template( - config={"interfaces": [{"name": "eth0", "type": "{{ interface_type }}"}]}, - default_values={"interface_type": "ethernet"}, - ) - config.templates.add(template) - old_checksum = config.get_cached_checksum() - - # Changing DeviceGroup.context should invalidate the cache - # and a new checksum should be returned in the following request. - device_group.context = {"interface_type": "virtual"} - device_group.full_clean() - device_group.save() - - # Config.backend_instance is a "cached_property", hence deleting - # the attribute is required for testing. - del config.backend_instance - - new_checksum = config.get_cached_checksum() - self.assertNotEqual(old_checksum, new_checksum) - def test_management_ip_changed_not_emitted_on_creation(self): with catch_signal(management_ip_changed) as handler: self._create_device(organization=self._get_org()) @@ -533,25 +482,6 @@ def test_manage_devices_group_templates_skips_deactivated_devices(self): # Status must remain "deactivated" — no config push is initiated. self.assertEqual(device.config.status, "deactivated") - def test_device_field_changed_checks(self): - self._create_device() - device_group = self._create_device_group() - with self.subTest("Deferred fields remained deferred"): - device = Device.objects.only("id", "created").first() - device._check_changed_fields() - - with self.subTest("Deferred fields becomes non-deferred"): - device.name = "new-name" - device.management_ip = "10.0.0.1" - device.group_id = device_group.id - device.organization_id = self._create_org().id - # assigning a random ip to last_ip - device.last_ip = "172.217.22.14" - # Another query is generated due to "config.set_status_modified" - # on name change - with self.assertNumQueries(3): - device._check_changed_fields() - @mock.patch.object(app_settings, "WHOIS_CONFIGURED", True) def test_changed_checked_fields_no_duplicates(self): """Ensure `_changed_checked_fields` contains `last_ip` only once. @@ -660,6 +590,80 @@ class TestTransactionDevice( CreateDeviceGroupMixin, TransactionTestCase, ): + def test_device_field_changed_checks(self): + self._create_device() + device_group = self._create_device_group() + with self.subTest("Deferred fields remained deferred"): + device = Device.objects.only("id", "created").first() + device._check_changed_fields() + + with self.subTest("Deferred fields becomes non-deferred"): + device.name = "new-name" + device.management_ip = "10.0.0.1" + device.group_id = device_group.id + device.organization_id = self._create_org().id + # assigning a random ip to last_ip + device.last_ip = "172.217.22.14" + # Changing group_id fires device_group_changed, which has two + # transaction.on_commit-deferred handlers: the devicegroup cache + # invalidation (the one this test guards) and the unrelated + # template-propagation task (devicegroup_templates_change_handler). + # Both run for real here because TransactionTestCase commits + # instead of rolling back, so their combined queries are counted. + with self.assertNumQueries(6): + device._check_changed_fields() + + def test_changing_org_variable_invalidates_cache(self): + org = self._get_org() + config_settings = OrganizationConfigSettings.objects.create( + organization=org, context={} + ) + config = self._create_config(organization=org) + template = self._create_template( + config={"interfaces": [{"name": "eth0", "type": "{{ interface_type }}"}]}, + default_values={"interface_type": "ethernet"}, + ) + config.templates.add(template) + old_checksum = config.get_cached_checksum() + + # Changing OrganizationConfigSettings.context should invalidate the cache + # and a new checksum should be returned in the following request. + config_settings.context = {"interface_type": "virtual"} + config_settings.full_clean() + config_settings.save() + + # Config.backend_instance is a "cached_property", hence deleting + # the attribute is required for testing. + del config.backend_instance + + new_checksum = config.get_cached_checksum() + self.assertNotEqual(old_checksum, new_checksum) + + def test_changing_group_variable_invalidates_cache(self): + org = self._get_org() + device_group = self._create_device_group(organization=org, context={}) + device = self._create_device(organization=org, group=device_group) + config = self._create_config(device=device) + template = self._create_template( + config={"interfaces": [{"name": "eth0", "type": "{{ interface_type }}"}]}, + default_values={"interface_type": "ethernet"}, + ) + config.templates.add(template) + old_checksum = config.get_cached_checksum() + + # Changing DeviceGroup.context should invalidate the cache + # and a new checksum should be returned in the following request. + device_group.context = {"interface_type": "virtual"} + device_group.full_clean() + device_group.save() + + # Config.backend_instance is a "cached_property", hence deleting + # the attribute is required for testing. + del config.backend_instance + + new_checksum = config.get_cached_checksum() + self.assertNotEqual(old_checksum, new_checksum) + def test_deactivating_device_with_config(self): self._create_template(required=True) self._create_template(name="Default", default=True) diff --git a/openwisp_controller/config/tests/test_vpn.py b/openwisp_controller/config/tests/test_vpn.py index 211ad4a06..e1611e740 100644 --- a/openwisp_controller/config/tests/test_vpn.py +++ b/openwisp_controller/config/tests/test_vpn.py @@ -20,7 +20,11 @@ from ..exceptions import ZeroTierIdentityGenerationError from ..settings import API_TASK_RETRY_OPTIONS from ..signals import config_modified, vpn_peers_changed, vpn_server_modified -from ..tasks import create_vpn_dh, trigger_vpn_server_endpoint +from ..tasks import ( + create_vpn_dh, + invalidate_vpn_server_devices_cache_change, + trigger_vpn_server_endpoint, +) from .utils import ( CreateConfigTemplateMixin, TestVpnX509Mixin, @@ -518,11 +522,10 @@ def test_update_vpn_dh(self, dhparam): def test_vpn_server_change_invalidates_device_cache(self): device, vpn, template = self._create_wireguard_vpn_template() - with catch_signal( - vpn_server_modified - ) as mocked_vpn_server_modified, catch_signal( - config_modified - ) as mocked_config_modified: + with ( + catch_signal(vpn_server_modified) as mocked_vpn_server_modified, + catch_signal(config_modified) as mocked_config_modified, + ): vpn.host = "localhost" vpn.save(update_fields=["host"]) mocked_vpn_server_modified.assert_called_once_with( @@ -538,6 +541,92 @@ def test_vpn_server_change_invalidates_device_cache(self): device=device, ) + def test_vpn_server_change_updates_client_checksum_db(self): + # the WireGuard client renders the VPN host as the peer + # "endpoint_host", so changing the VPN host alters the client + # configuration + device, vpn, _ = self._create_wireguard_vpn_template() + config = Config.objects.get(pk=device.config.pk) + old_checksum_db = config.checksum_db + # sanity check: the stored checksum initially matches the + # freshly computed checksum for the client configuration + self.assertEqual(old_checksum_db, config.checksum) + # change a VPN server field that is part of the client configuration + vpn.host = "changed.example.com" + # select_related on the client queryset keeps this bounded: dropping + # it reintroduces the per-client N+1 and increases the query count. + with self.assertNumQueries(25): + vpn.save(update_fields=["host"]) + config = Config.objects.get(pk=device.config.pk) + # the client's stored checksum must reflect the new VPN server context + self.assertNotEqual(config.checksum_db, old_checksum_db) + self.assertEqual(config.checksum_db, config.checksum) + self.assertEqual(config.status, "modified") + + def test_invalidate_clients_cache_no_checksum_change(self): + # config_modified must not be emitted for a client whose checksum + # did not actually change (unlike the old, unconditional behavior) + device, vpn, _ = self._create_wireguard_vpn_template() + config = Config.objects.get(pk=device.config.pk) + old_checksum_db = config.checksum_db + with catch_signal(config_modified) as mocked_config_modified: + # select_related on the client queryset keeps this bounded: dropping + # it reintroduces the per-client N+1 and increases the query count. + with self.assertNumQueries(19): + VpnClient.invalidate_clients_cache(vpn) + mocked_config_modified.assert_not_called() + config = Config.objects.get(pk=device.config.pk) + self.assertEqual(config.checksum_db, old_checksum_db) + + def test_ca_renew_invalidates_vpn_checksum(self): + vpn = self._create_vpn() + with catch_signal(vpn_server_modified) as mocked: + vpn.ca.renew() + # vpn_server_modified fires via the CacheDependency, + # which cascades to client config invalidation + self.assertTrue(mocked.called) + + def test_cert_renew_invalidates_vpn_checksum(self): + vpn = self._create_vpn() + with catch_signal(vpn_server_modified) as mocked: + vpn.cert.renew() + self.assertTrue(mocked.called) + + def _setup_vpn_client_config(self): + """ + Creates a VPN server, a VPN template and a client device/config using + it, then returns ``(vpn, config, old_checksum_db)`` with the client + config's stored checksum captured and asserted to be up to date. + """ + vpn = self._create_vpn() + vpn_template = self._create_template( + name="vpn-template", type="vpn", vpn=vpn, config={} + ) + device = self._create_device_config() + device.config.templates.add(vpn_template) + config = Config.objects.get(pk=device.config.pk) + old_checksum_db = config.checksum_db + self.assertEqual(old_checksum_db, config.checksum) + return vpn, config, old_checksum_db + + def test_ca_renew_cascades_to_client_config(self): + vpn, config, old_checksum_db = self._setup_vpn_client_config() + vpn.ca.renew() + config = Config.objects.get(pk=config.pk) + self.assertNotEqual(config.checksum_db, old_checksum_db) + self.assertEqual(config.checksum_db, config.checksum) + self.assertEqual(config.status, "modified") + + def test_cert_renew_cascades_to_client_config(self): + # Only the signal is asserted here, unlike the CA-renew counterpart: + # a client's OpenVPN config embeds the server's certificate, not + # the server's own certificate, so renewing the server cert does not + # change the client's rendered config or its checksum. + vpn, _, _ = self._setup_vpn_client_config() + with catch_signal(vpn_server_modified) as mocked_vpn_server_modified: + vpn.cert.renew() + mocked_vpn_server_modified.assert_called_once() + class TestWireguard(BaseTestVpn, TestWireguardVpnMixin, TestCase): def test_wireguard_config_creation(self): @@ -837,10 +926,11 @@ def test_update_vpn_server_configuration(self): success_response.status_code = 200 success_response.raise_for_status = mock.Mock() - with mock.patch( - "openwisp_controller.config.tasks.logger.info" - ) as mocked_logger, mock.patch( - "requests.post", return_value=success_response + with ( + mock.patch( + "openwisp_controller.config.tasks.logger.info" + ) as mocked_logger, + mock.patch("requests.post", return_value=success_response), ): vpn.save() vpn_client.refresh_from_db() @@ -858,8 +948,9 @@ def test_update_vpn_server_configuration(self): fail_response.raise_for_status.side_effect = requests.exceptions.HTTPError( "Not Found" ) - with mock.patch("logging.Logger.warning") as mocked_logger, mock.patch( - "requests.post", return_value=fail_response + with ( + mock.patch("logging.Logger.warning") as mocked_logger, + mock.patch("requests.post", return_value=fail_response), ): post_save.send( instance=vpn_client, sender=vpn_client._meta.model, created=False @@ -1472,6 +1563,80 @@ def test_zerotier_change_vpn_backend_with_vpnclient(self): context_manager.exception.message_dict, expected_error_dict ) + @mock.patch(_ZT_GENERATE_IDENTITY_SUBPROCESS) + @mock.patch(_ZT_SERVICE_REQUESTS) + def test_zerotier_vpn_name_change_sends_signal( + self, mock_requests, mock_subprocess + ): + mock_requests.get.side_effect = [ + self._get_mock_response(200, response=self._TEST_ZT_NODE_CONFIG) + ] + mock_requests.post.side_effect = [ + self._get_mock_response(200), + self._get_mock_response(200), + self._get_mock_response(200), + self._get_mock_response(200), + ] + self._set_subprocess_mock(mock_subprocess) + device, vpn, template = self._create_zerotier_vpn_template() + + with catch_signal(vpn_server_modified) as handler: + vpn.name = "updated-zerotier-vpn" + vpn.save() + + handler.assert_called_once() + self.assertEqual(handler.call_args[1]["instance"].pk, vpn.pk) + + @mock.patch(_ZT_GENERATE_IDENTITY_SUBPROCESS) + @mock.patch(_ZT_SERVICE_REQUESTS) + def test_zerotier_vpn_name_change_updates_client_checksum( + self, mock_requests, mock_subprocess + ): + mock_requests.get.side_effect = [ + self._get_mock_response(200, response=self._TEST_ZT_NODE_CONFIG) + ] + mock_requests.post.side_effect = [ + self._get_mock_response(200), + self._get_mock_response(200), + self._get_mock_response(200), + self._get_mock_response(200), + ] + self._set_subprocess_mock(mock_subprocess) + device, vpn, template = self._create_zerotier_vpn_template() + config = device.config + pk = vpn.pk.hex + # Set up a device config that references the VPN name context variable + # so the checksum depends on Vpn.name + config.config = { + "files": [ + { + "path": "/tmp/test", + "mode": "0644", + "contents": "{{ network_name_" + pk + " }}", + } + ] + } + config.save() + old_checksum_db = config.checksum_db + with mock.patch( + "openwisp_controller.config.tasks" + ".invalidate_vpn_server_devices_cache_change.delay" + ) as mocked_delay: + with self.captureOnCommitCallbacks(execute=True): + vpn.name = "updated-zerotier-vpn" + vpn.save() + mocked_delay.assert_not_called() + mocked_delay.assert_called_once_with(vpn.id) + + # Trigger the actual cache invalidation to verify checksum changes + invalidate_vpn_server_devices_cache_change(vpn.pk) + config.refresh_from_db() + # Invalidate cached backend_instance so checksum recomputes fresh + config._invalidate_backend_instance_cache() + self.assertNotEqual(config.checksum_db, old_checksum_db) + self.assertEqual(config.checksum_db, config.checksum) + self.assertEqual(config.status, "modified") + class TestZeroTierTransaction( BaseTestVpn, TestZeroTierVpnMixin, TestWireguardVpnMixin, TransactionTestCase @@ -1873,18 +2038,22 @@ def test_zerotier_update_vpn_server_configuration( mock_error.reset_mock() mock_requests.reset_mock() - with self.subTest( - "Test zerotier configuration update " - "with retry mechanism (recoverable errors)" - ), mock.patch("celery.app.task.Task.request") as mock_task_request: + with ( + self.subTest( + "Test zerotier configuration update " + "with retry mechanism (recoverable errors)" + ), + mock.patch("celery.app.task.Task.request") as mock_task_request, + ): max_retries = API_TASK_RETRY_OPTIONS.get("max_retries") mock_task_request.called_directly = False config = vpn.get_config()["zerotier"][0] config.update({"private": True}) - with self.subTest( - "Test update when max retry limit is not reached" - ), self.assertRaises(Retry): + with ( + self.subTest("Test update when max retry limit is not reached"), + self.assertRaises(Retry), + ): mock_requests.get.side_effect = [ # For node status self._get_mock_response(200, response=self._TEST_ZT_NODE_CONFIG) @@ -1935,9 +2104,10 @@ def test_zerotier_update_vpn_server_configuration( # During the last attempt, the task will give up # retrying and raise a 'RequestException', # which will be handled and logged as an error - with self.subTest( - "Test update when max retry limit is reached" - ), self.assertRaises(RequestException): + with ( + self.subTest("Test update when max retry limit is reached"), + self.assertRaises(RequestException), + ): mock_requests.get.side_effect = [ # For node status self._get_mock_response(200, response=self._TEST_ZT_NODE_CONFIG) diff --git a/openwisp_controller/config/tests/utils.py b/openwisp_controller/config/tests/utils.py index aa8b2d910..109788f9d 100644 --- a/openwisp_controller/config/tests/utils.py +++ b/openwisp_controller/config/tests/utils.py @@ -177,6 +177,7 @@ def _create_wireguard_vpn_template( vpn=vpn, organization=org1, auto_cert=auto_cert, + config={}, ) device = self._create_device_config() device.config.templates.add(template) diff --git a/openwisp_controller/pki/tests/test_api.py b/openwisp_controller/pki/tests/test_api.py index 192562d46..51faebdaf 100644 --- a/openwisp_controller/pki/tests/test_api.py +++ b/openwisp_controller/pki/tests/test_api.py @@ -1,4 +1,4 @@ -from django.test import TestCase +from django.test import TestCase, TransactionTestCase from django.urls import reverse from packaging.version import parse as parse_version from rest_framework import VERSION as REST_FRAMEWORK_VERSION @@ -161,7 +161,7 @@ def test_ca_post_renew_api(self): ca1 = self._create_ca(name="ca1", organization=self._get_org()) old_serial_num = ca1.serial_number path = reverse("pki_api:ca_renew", args=[ca1.pk]) - with self.assertNumQueries(4): + with self.assertNumQueries(5): r = self.client.post(path) ca1.refresh_from_db() self.assertEqual(r.status_code, 200) @@ -244,35 +244,10 @@ def test_cert_detail_api(self): self.assertEqual(r.data["id"], cert1.pk) self.assertEqual(r.data["extensions"], []) - def test_cert_put_api(self): - cert1 = self._create_cert(name="cert1") - org2 = self._create_org() - path = reverse("pki_api:cert_detail", args=[cert1.pk]) - data = { - "name": "cert1-change", - "organization": org2.pk, - "notes": "new-notes", - } - with self.assertNumQueries(10): - r = self.client.put(path, data, content_type="application/json") - self.assertEqual(r.status_code, 200) - self.assertEqual(r.data["name"], "cert1-change") - self.assertEqual(r.data["organization"], org2.pk) - self.assertEqual(r.data["notes"], "new-notes") - - def test_cert_patch_api(self): - cert1 = self._create_cert(name="cert1") - path = reverse("pki_api:cert_detail", args=[cert1.pk]) - data = {"name": "cert1-change"} - with self.assertNumQueries(8): - r = self.client.patch(path, data, content_type="application/json") - self.assertEqual(r.status_code, 200) - self.assertEqual(r.data["name"], "cert1-change") - def test_cert_delete_api(self): cert1 = self._create_cert(name="cert1") path = reverse("pki_api:cert_detail", args=[cert1.pk]) - with self.assertNumQueries(6): + with self.assertNumQueries(5): r = self.client.delete(path) self.assertEqual(r.status_code, 204) self.assertEqual(Cert.objects.count(), 0) @@ -285,28 +260,6 @@ def test_ca_in_cert_detail_fields(self): self.assertEqual(r.status_code, 200) self.assertEqual(r.data["ca"], cert1.ca.id) - def test_post_cert_renew_api(self): - cert1 = self._create_cert(name="cert1") - old_serial_num = cert1.serial_number - path = reverse("pki_api:cert_renew", args=[cert1.pk]) - with self.assertNumQueries(5): - r = self.client.post(path) - self.assertEqual(r.status_code, 200) - cert1.refresh_from_db() - self.assertNotEqual(cert1.serial_number, old_serial_num) - self.assertNotEqual(r.data["serial_number"], old_serial_num) - - def test_post_cert_revoke_api(self): - cert1 = self._create_cert(name="cert1") - self.assertFalse(cert1.revoked) - path = reverse("pki_api:cert_revoke", args=[cert1.pk]) - with self.assertNumQueries(4): - r = self.client.post(path) - cert1.refresh_from_db() - self.assertEqual(r.status_code, 200) - self.assertTrue(cert1.revoked) - self.assertTrue(r.data["revoked"]) - @capture_any_output() def test_bearer_authentication(self): self.client.logout() @@ -368,3 +321,63 @@ def test_bearer_authentication(self): HTTP_AUTHORIZATION=f"Bearer {token}", ) self.assertEqual(response.status_code, 200) + + +class TestTransactionPkiApi( + AssertNumQueriesSubTestMixin, + TestAdminMixin, + TestPkiMixin, + TestOrganizationMixin, + AuthenticationMixin, + TransactionTestCase, +): + def setUp(self): + super().setUp() + self._login() + + def test_cert_put_api(self): + cert1 = self._create_cert(name="cert1") + org2 = self._create_org() + path = reverse("pki_api:cert_detail", args=[cert1.pk]) + data = { + "name": "cert1-change", + "organization": org2.pk, + "notes": "new-notes", + } + with self.assertNumQueries(10): + r = self.client.put(path, data, content_type="application/json") + self.assertEqual(r.status_code, 200) + self.assertEqual(r.data["name"], "cert1-change") + self.assertEqual(r.data["organization"], org2.pk) + self.assertEqual(r.data["notes"], "new-notes") + + def test_cert_patch_api(self): + cert1 = self._create_cert(name="cert1") + path = reverse("pki_api:cert_detail", args=[cert1.pk]) + data = {"name": "cert1-change"} + with self.assertNumQueries(8): + r = self.client.patch(path, data, content_type="application/json") + self.assertEqual(r.status_code, 200) + self.assertEqual(r.data["name"], "cert1-change") + + def test_post_cert_renew_api(self): + cert1 = self._create_cert(name="cert1") + old_serial_num = cert1.serial_number + path = reverse("pki_api:cert_renew", args=[cert1.pk]) + with self.assertNumQueries(6): + r = self.client.post(path) + self.assertEqual(r.status_code, 200) + cert1.refresh_from_db() + self.assertNotEqual(cert1.serial_number, old_serial_num) + self.assertNotEqual(r.data["serial_number"], old_serial_num) + + def test_post_cert_revoke_api(self): + cert1 = self._create_cert(name="cert1") + self.assertFalse(cert1.revoked) + path = reverse("pki_api:cert_revoke", args=[cert1.pk]) + with self.assertNumQueries(4): + r = self.client.post(path) + cert1.refresh_from_db() + self.assertEqual(r.status_code, 200) + self.assertTrue(cert1.revoked) + self.assertTrue(r.data["revoked"]) diff --git a/openwisp_controller/subnet_division/admin.py b/openwisp_controller/subnet_division/admin.py index 287cc0c28..d1a1eb121 100644 --- a/openwisp_controller/subnet_division/admin.py +++ b/openwisp_controller/subnet_division/admin.py @@ -28,7 +28,7 @@ class SubnetDivisionRuleInlineAdmin( help_text = { "text": _( "Please keep in mind that once the subnet division rule is created " - 'changing changing "Size", "Number of Subnets" or decreasing ' + 'changing "Size", "Number of Subnets" or decreasing ' '"Number of IPs" will not be possible.' ), "documentation_url": ( diff --git a/openwisp_controller/subnet_division/rule_types/vpn.py b/openwisp_controller/subnet_division/rule_types/vpn.py index 2680f89fb..8e0260e42 100644 --- a/openwisp_controller/subnet_division/rule_types/vpn.py +++ b/openwisp_controller/subnet_division/rule_types/vpn.py @@ -43,9 +43,11 @@ def provision_for_existing_objects(cls, rule_obj): @classmethod def post_provision_handler(cls, instance, provisioned, **kwargs): - super().post_provision_handler(instance, provisioned, **kwargs) # Assign the first provisioned IP address to the VPNClient - # only when subnets and IPs have been provisioned + # only when subnets and IPs have been provisioned, and do so + # *before* calling the superclass handler: it computes and caches + # the Config checksum, which needs instance.ip already set so + # the "ip_address_" template variable is resolved. if provisioned and provisioned["ip_addresses"]: # Delete any previously assigned IP address if instance.ip: @@ -53,6 +55,7 @@ def post_provision_handler(cls, instance, provisioned, **kwargs): instance.ip = provisioned["ip_addresses"][0] instance.full_clean() instance.save() + super().post_provision_handler(instance, provisioned, **kwargs) @classmethod def destroy_provisioned_subnets_ips(cls, instance, **kwargs): diff --git a/openwisp_controller/subnet_division/tasks.py b/openwisp_controller/subnet_division/tasks.py index 49222e9a3..998e9a1c2 100644 --- a/openwisp_controller/subnet_division/tasks.py +++ b/openwisp_controller/subnet_division/tasks.py @@ -38,6 +38,13 @@ def update_subnet_division_index(rule_id): ) index.save() + config_ids = ( + division_rule.subnetdivisionindex_set.filter(config_id__isnull=False) + .values_list("config_id", flat=True) + .distinct() + ) + Config.bulk_invalidate_get_cached_checksum({"id__in": list(config_ids)}) + @shared_task def update_subnet_name_description(rule_id): @@ -74,6 +81,8 @@ def provision_extra_ips(rule_id, old_number_of_ips): def _create_ipaddress_and_subnetdivision_index_objects(ips, indexes): IpAddress.objects.bulk_create(ips) SubnetDivisionIndex.objects.bulk_create(indexes) + config_ids = {index.config_id for index in indexes} + Config.bulk_invalidate_get_cached_checksum({"id__in": config_ids}) generated_ips = [] generated_indexes = [] diff --git a/openwisp_controller/subnet_division/tests/test_models.py b/openwisp_controller/subnet_division/tests/test_models.py index 2377399c8..071fb576f 100644 --- a/openwisp_controller/subnet_division/tests/test_models.py +++ b/openwisp_controller/subnet_division/tests/test_models.py @@ -23,6 +23,7 @@ IpAddress = load_model("openwisp_ipam", "IpAddress") SubnetDivisionRule = load_model("subnet_division", "SubnetDivisionRule") SubnetDivisionIndex = load_model("subnet_division", "SubnetDivisionIndex") +Config = load_model("config", "Config") VpnClient = load_model("config", "VpnClient") Device = load_model("config", "Device") OrganizationConfigSettings = load_model("config", "OrganizationConfigSettings") @@ -62,6 +63,35 @@ def test_provisioned_subnets(self): for ip_id in range(1, rule.number_of_ips + 1): self.assertIn(f"{rule.label}_subnet{subnet_id}_ip{ip_id}", context) + def test_vpn_client_ip_resolved_before_checksum_cached(self): + """ + Regression test: VpnSubnetDivisionRuleType.post_provision_handler + used to call super().post_provision_handler() (which computes and + caches the Config checksum) before assigning the provisioned IP + address to the VpnClient. This left the "ip_address_" + template variable unresolved in the cached checksum. + """ + self._get_vpn_subdivision_rule() + # "config={}" makes Template.clean() auto-generate the wireguard + # client config via Vpn.auto_client(), which is what embeds the + # "{{ip_address_}}" placeholder that must be resolved. + wireguard_template = self._create_template( + name="wireguard-auto-client", + type="vpn", + vpn=self.vpn_server, + organization=self.org, + auto_cert=True, + config={}, + ) + self.config.templates.add(wireguard_template) + vpnclient = self.config.vpnclient_set.get(vpn=self.vpn_server) + self.assertNotEqual(vpnclient.ip, None) + ip_address_key = f"ip_address_{self.vpn_server.pk.hex}" + context = self.config.get_context() + self.assertEqual(context[ip_address_key], vpnclient.ip.ip_address) + refreshed_config = Config.objects.get(pk=self.config.pk) + self.assertEqual(refreshed_config.checksum_db, refreshed_config.checksum) + class TestSubnetDivisionRule( SubnetDivisionTestMixin, @@ -333,8 +363,13 @@ def test_rule_label_updated(self): ) index_count = index_queryset.count() subnet_count = subnet_queryset.count() - rule.label = new_rule_label - rule.save() + with patch.object(Config, "bulk_invalidate_get_cached_checksum") as mocked: + rule.label = new_rule_label + rule.save() + # Regression test: renaming a rule's label rewrites the keywords of + # its SubnetDivisionIndex entries, which feed Config.get_context(); + # the affected configs' checksums must be invalidated accordingly. + mocked.assert_called_once_with({"id__in": [self.config.id]}) rule.refresh_from_db() self.assertEqual(rule.label, new_rule_label) @@ -375,8 +410,13 @@ def test_number_of_ips_updated(self): ) new_number_of_ips = rule.number_of_ips + 2 - rule.number_of_ips = new_number_of_ips - rule.save() + with patch.object(Config, "bulk_invalidate_get_cached_checksum") as mocked: + rule.number_of_ips = new_number_of_ips + rule.save() + # Regression test: provisioning extra IPs creates new + # SubnetDivisionIndex entries, which feed Config.get_context(); + # the affected configs' checksums must be invalidated accordingly. + mocked.assert_called_once_with({"id__in": {self.config.id}}) rule.refresh_from_db() self.assertEqual(rule.number_of_ips, new_number_of_ips)