Просмотр исходного кода

perf(api): Prefetch only for requested fields

Device and interface viewsets declare the prefetches their method and
property fields need. BaseViewSet attaches only those of the response
fields, and all of them for exports and writes. The new
get_fields_for_serializer() and response_fields share the field
selection logic.

Fixes #23292
Martin Hauser 2 дней назад
Родитель
Сommit
85fa431fc6
4 измененных файлов с 394 добавлено и 44 удалено
  1. 23 23
      netbox/dcim/api/views.py
  2. 327 5
      netbox/dcim/tests/test_api.py
  3. 30 1
      netbox/netbox/api/viewsets/__init__.py
  4. 14 15
      netbox/utilities/api.py

+ 23 - 23
netbox/dcim/api/views.py

@@ -43,7 +43,9 @@ class PathEndpointMixin:
         """
         """
         Trace a complete cable path and return each segment as a three-tuple of (termination, cable, termination).
         Trace a complete cable path and return each segment as a three-tuple of (termination, cable, termination).
         """
         """
-        obj = get_object_or_404(self.queryset, pk=pk)
+        # Path nodes load as they do for the connected endpoint fields
+        queryset = self.queryset.prefetch_related(*self.get_field_prefetches({'connected_endpoints'}))
+        obj = get_object_or_404(queryset, pk=pk)
 
 
         # Initialize the path array
         # Initialize the path array
         path = []
         path = []
@@ -426,13 +428,15 @@ class PlatformViewSet(NetBoxModelViewSet):
 #
 #
 
 
 class DeviceViewSet(ConfigContextQuerySetMixin, RenderConfigMixin, NetBoxModelViewSet):
 class DeviceViewSet(ConfigContextQuerySetMixin, RenderConfigMixin, NetBoxModelViewSet):
-    queryset = Device.objects.prefetch_related(
-        'device_type__manufacturer',  # Referenced by Device.__str__() for unnamed devices
-        'parent_bay',  # Referenced by DeviceSerializer.get_parent_device()
-    )
+    queryset = Device.objects.all()
     serializer_class = serializers.DeviceSerializer
     serializer_class = serializers.DeviceSerializer
     filterset_class = filtersets.DeviceFilterSet
     filterset_class = filtersets.DeviceFilterSet
     pagination_class = StripCountAnnotationsPaginator
     pagination_class = StripCountAnnotationsPaginator
+    field_prefetches = {
+        # Read by Device.__str__() for unnamed devices and by config context rendered without its annotation
+        ('display', 'config_context'): ('device_type__manufacturer',),
+        ('parent_device',): ('parent_bay',),
+    }
 
 
 
 
 class VirtualDeviceContextViewSet(NetBoxModelViewSet):
 class VirtualDeviceContextViewSet(NetBoxModelViewSet):
@@ -499,26 +503,22 @@ class CoolingOutflowViewSet(NetBoxModelViewSet):
 
 
 
 
 class InterfaceViewSet(PathEndpointMixin, NetBoxModelViewSet):
 class InterfaceViewSet(PathEndpointMixin, NetBoxModelViewSet):
-    queryset = Interface.objects.prefetch_related(
-        GenericPrefetch(
-            "cable__terminations__termination",
-            [
-                Interface.objects.select_related("device", "cable"),
-            ],
-        ),
-        GenericPrefetch(
-            "_path__path_objects",
-            [
-                Interface.objects.select_related("device", "cable"),
-            ],
-        ),
-        'virtual_circuit_termination',
-        'l2vpn_terminations',  # Referenced by InterfaceSerializer.l2vpn_termination
-        'ip_addresses',  # Referenced by Interface.count_ipaddresses()
-        'fhrp_group_assignments',  # Referenced by Interface.count_fhrp_groups()
-    )
+    queryset = Interface.objects.all()
     serializer_class = serializers.InterfaceSerializer
     serializer_class = serializers.InterfaceSerializer
     filterset_class = filtersets.InterfaceFilterSet
     filterset_class = filtersets.InterfaceFilterSet
+    field_prefetches = {
+        ('link_peers', 'link_peers_type'): (
+            GenericPrefetch('cable__terminations__termination', [Interface.objects.select_related('device', 'cable')]),
+        ),
+        ('connected_endpoints', 'connected_endpoints_type'): (
+            GenericPrefetch('_path__path_objects', [Interface.objects.select_related('device', 'cable')]),
+            'virtual_circuit_termination',
+        ),
+        ('connected_endpoints_reachable',): ('_path',),
+        ('l2vpn_termination',): ('l2vpn_terminations',),
+        ('count_ipaddresses',): ('ip_addresses',),
+        ('count_fhrp_groups',): ('fhrp_group_assignments',),
+    }
 
 
     def get_bulk_destroy_queryset(self):
     def get_bulk_destroy_queryset(self):
         # Ensure child interfaces are deleted prior to their parents
         # Ensure child interfaces are deleted prior to their parents

+ 327 - 5
netbox/dcim/tests/test_api.py

@@ -1,28 +1,34 @@
 import json
 import json
 import os
 import os
 import tempfile
 import tempfile
+from contextlib import contextmanager
 from unittest.mock import patch
 from unittest.mock import patch
 
 
 from django.conf import settings
 from django.conf import settings
 from django.contrib.contenttypes.models import ContentType
 from django.contrib.contenttypes.models import ContentType
+from django.contrib.contenttypes.prefetch import GenericPrefetch
 from django.db import connection
 from django.db import connection
 from django.test import override_settings, tag
 from django.test import override_settings, tag
 from django.test.utils import CaptureQueriesContext
 from django.test.utils import CaptureQueriesContext
 from django.urls import reverse
 from django.urls import reverse
 from django.utils.translation import gettext as _
 from django.utils.translation import gettext as _
 from rest_framework import status
 from rest_framework import status
+from rest_framework.test import APIRequestFactory
 
 
+from circuits.choices import VirtualCircuitTerminationRoleChoices
+from circuits.models import Provider, ProviderNetwork, VirtualCircuit, VirtualCircuitTermination, VirtualCircuitType
 from core.choices import ObjectChangeActionChoices
 from core.choices import ObjectChangeActionChoices
 from core.models import ObjectChange, ObjectType
 from core.models import ObjectChange, ObjectType
 from dcim.api.serializers import InterfaceSerializer
 from dcim.api.serializers import InterfaceSerializer
 from dcim.api.serializers_.nested import NestedDeviceBaySerializer, NestedDeviceSerializer
 from dcim.api.serializers_.nested import NestedDeviceBaySerializer, NestedDeviceSerializer
+from dcim.api.views import DeviceViewSet, InterfaceViewSet
 from dcim.choices import *
 from dcim.choices import *
 from dcim.constants import *
 from dcim.constants import *
 from dcim.graphql.types import _CABLE_TERMINATION_MODELS
 from dcim.graphql.types import _CABLE_TERMINATION_MODELS
 from dcim.models import *
 from dcim.models import *
-from extras.models import ConfigTemplate, Tag
-from ipam.choices import VLANQinQRoleChoices
-from ipam.models import ASN, RIR, VLAN, VRF, IPAddress
+from extras.models import ConfigTemplate, ExportTemplate, Tag
+from ipam.choices import FHRPGroupProtocolChoices, VLANQinQRoleChoices
+from ipam.models import ASN, RIR, VLAN, VRF, FHRPGroup, FHRPGroupAssignment, IPAddress
 from netbox.api.serializers import GenericObjectSerializer
 from netbox.api.serializers import GenericObjectSerializer
 from tenancy.models import Tenant
 from tenancy.models import Tenant
 from users.constants import TOKEN_PREFIX
 from users.constants import TOKEN_PREFIX
@@ -39,6 +45,8 @@ from utilities.testing import (
     disable_warnings,
     disable_warnings,
 )
 )
 from virtualization.models import Cluster, ClusterType
 from virtualization.models import Cluster, ClusterType
+from vpn.choices import L2VPNTypeChoices
+from vpn.models import L2VPN, L2VPNTermination
 from wireless.choices import WirelessChannelChoices
 from wireless.choices import WirelessChannelChoices
 from wireless.models import WirelessLAN
 from wireless.models import WirelessLAN
 
 
@@ -89,6 +97,100 @@ class Mixins:
             self.assertEqual(segment1[1]['label'], cable.label)
             self.assertEqual(segment1[1]['label'], cable.label)
             self.assertEqual(segment1[2][0]['name'], peer_obj.name)
             self.assertEqual(segment1[2][0]['name'], peer_obj.name)
 
 
+    class FieldPrefetchMixin(APITestCase):
+        viewset = None
+        # Narrowed requests compared in addition to the default, brief and single-field ones
+        field_prefetch_cases = ()
+        # Literal reference lookups (the static prefetches before the declarations) loaded unconditionally
+        eager_prefetches = ()
+
+        @staticmethod
+        def lookup_path(lookup):
+            return getattr(lookup, 'prefetch_to', lookup)
+
+        @contextmanager
+        def capture_querysets(self, viewset):
+            """Collect the querysets returned by the viewset's get_queryset()."""
+            querysets = []
+            get_queryset = viewset.get_queryset
+
+            def spy(view):
+                querysets.append(get_queryset(view))
+                return querysets[-1]
+
+            with patch.object(viewset, 'get_queryset', spy):
+                yield querysets
+
+        def assert_lookups(self, querysets, expected):
+            """Assert each queryset leads with exactly these lookups, as these objects, and no other declared one."""
+            declared = [lookup for lookups in self.viewset.field_prefetches.values() for lookup in lookups]
+            self.assertTrue(querysets)
+            for queryset in querysets:
+                lookups = queryset._prefetch_related_lookups
+                paths = [self.lookup_path(lookup) for lookup in lookups]
+                self.assertEqual(paths[:len(expected)], [self.lookup_path(lookup) for lookup in expected], paths)
+                self.assertTrue(all(a is b for a, b in zip(lookups, expected)), paths)
+                tail = [lookup for lookup in lookups[len(expected):] if any(lookup is d for d in declared)]
+                self.assertFalse(tail, paths)
+
+        def test_list_objects_field_prefetches(self):
+            """Narrowed responses keep the reference content within its queries, and brief ones skip unread lookups."""
+            field_prefetches = self.viewset.field_prefetches
+            declared = [lookup for lookups in field_prefetches.values() for lookup in lookups]
+            fields = self.viewset.serializer_class.Meta.fields
+            # The map names serializer fields only, and its lookups no longer sit on the class queryset
+            self.assertLessEqual(set().union(*field_prefetches), set(fields))
+            static = {self.lookup_path(lookup) for lookup in self.viewset.queryset._prefetch_related_lookups}
+            self.assertFalse(static.intersection(self.lookup_path(lookup) for lookup in declared))
+
+            self.add_permissions(f'{self.model._meta.app_label}.view_{self.model._meta.model_name}')
+            scopes = self._create_field_prefetch_scopes()
+            url = self._get_list_url()
+            # Runtime bound: skip fields named after model columns or forward relations that start no lookup
+            heads = {self.lookup_path(lookup).split('__')[0] for lookup in (*self.eager_prefetches, *declared)}
+            skipped = {
+                field.name for field in self.model._meta.get_fields()
+                if (field.concrete or (field.many_to_many and not field.auto_created)) and field.name not in heads
+            }
+            cases = (
+                '', 'brief=1', *self.field_prefetch_cases,
+                *(f'fields=id,{name}' for name in fields if name not in skipped),
+            )
+
+            def get(query, eager):
+                queryset, prefetches = self.viewset.queryset, field_prefetches
+                if eager:
+                    queryset, prefetches = queryset.prefetch_related(*self.eager_prefetches), {}
+                with (
+                    patch.object(self.viewset, 'queryset', queryset),
+                    patch.object(self.viewset, 'field_prefetches', prefetches),
+                    CaptureQueriesContext(connection) as queries,
+                ):
+                    response = self.client.get(f'{url}?{query}', **self.header)
+                # Token authentication refreshes last_used once a minute
+                count = sum('"users_token"' not in entry['sql'] for entry in queries.captured_queries)
+                self.assertHttpStatus(response, status.HTTP_200_OK)
+                return response.json(), count
+
+            # Warm per-process caches (e.g. ContentType) for every related type
+            self.client.get(f'{url}?{scopes[-1]}', **self.header)
+            for params in cases:
+                with self.subTest(params=params):
+                    # Some cases take paths of their own, such as config context without its annotation
+                    self.client.get(f'{url}?{scopes[-1]}&{params}', **self.header)
+                    savings = set()
+                    for scope in scopes:
+                        expected, eager_count = get(f'{scope}&{params}', eager=True)
+                        data, count = get(f'{scope}&{params}', eager=False)
+                        self.assertEqual(data, expected)
+                        self.assertLessEqual(count, eager_count)
+                        savings.add(eager_count - count)
+                    self.assertEqual(len(savings), 1, f'A field reads a lookup it does not declare: {savings}')
+                    if not params:
+                        self.assertEqual(savings, {0})
+                    elif params == 'brief=1':
+                        self.assertGreater(savings.pop(), 0)
+
 
 
 class RegionTestCase(APIViewTestCases.APIViewTestCase):
 class RegionTestCase(APIViewTestCases.APIViewTestCase):
     model = Region
     model = Region
@@ -2543,8 +2645,14 @@ class PlatformTestCase(APIViewTestCases.APIViewTestCase):
             platform.save()
             platform.save()
 
 
 
 
-class DeviceTestCase(APIViewTestCases.APIViewTestCase):
+class DeviceTestCase(Mixins.FieldPrefetchMixin, APIViewTestCases.APIViewTestCase):
     model = Device
     model = Device
+    viewset = DeviceViewSet
+    field_prefetch_cases = (
+        'omit=parent_device', 'fields=id,parent_device&omit=parent_device', 'brief=1&omit=display',
+        'brief=1&fields=id,config_context',
+    )
+    eager_prefetches = ('device_type__manufacturer', 'parent_bay')
     brief_fields = ['description', 'display', 'id', 'name', 'url']
     brief_fields = ['description', 'display', 'id', 'name', 'url']
     bulk_update_data = {
     bulk_update_data = {
         'status': 'failed',
         'status': 'failed',
@@ -2652,6 +2760,28 @@ class DeviceTestCase(APIViewTestCases.APIViewTestCase):
             },
             },
         ]
         ]
 
 
+    def _create_field_prefetch_scopes(self):
+        """Return filters for one and for two sites of parent, installed child and unnamed devices."""
+        device = Device.objects.get(name='Device 1')
+        parent_type = DeviceType.objects.create(
+            manufacturer=device.device_type.manufacturer, model='Unit Parent Type', slug='unit-parent-type',
+            subdevice_role=SubdeviceRoleChoices.ROLE_PARENT,
+        )
+        child_type = DeviceType.objects.create(
+            manufacturer=device.device_type.manufacturer, model='Unit Child Type', slug='unit-child-type',
+            subdevice_role=SubdeviceRoleChoices.ROLE_CHILD, u_height=0,
+        )
+        filters = []
+        for i in range(2):
+            site = Site.objects.create(name=f'Unit Site {i}', slug=f'unit-site-{i}')
+            parent = Device.objects.create(device_type=parent_type, role=device.role, site=site, name=f'Parent {i}')
+            child = Device.objects.create(device_type=child_type, role=device.role, site=site, name=f'Child {i}')
+            DeviceBay.objects.create(device=parent, name='Bay 1', installed_device=child)
+            # Unnamed, so its display reads the device type and manufacturer
+            create_test_device(None, site=site)
+            filters.append(f'site_id={site.pk}')
+        return [filters[0], '&'.join(filters)]
+
     def test_config_context_included_by_default_in_list_view(self):
     def test_config_context_included_by_default_in_list_view(self):
         """
         """
         Check that config context data is included by default in the devices list.
         Check that config context data is included by default in the devices list.
@@ -3074,6 +3204,33 @@ class DeviceTestCase(APIViewTestCases.APIViewTestCase):
             'No objects should be created when any sibling is aborted',
             'No objects should be created when any sibling is aborted',
         )
         )
 
 
+    def test_field_prefetches_follow_requested_fields(self):
+        """List and detail responses carry the declared prefetches of exactly the fields they include."""
+        self.add_permissions('dcim.view_device')
+        manufacturer = list(DeviceViewSet.field_prefetches[('display', 'config_context')])
+        parent = list(DeviceViewSet.field_prefetches[('parent_device',)])
+        # fields takes precedence over omit, and both over brief
+        cases = {
+            '': manufacturer + parent,
+            'brief=1': manufacturer,
+            'fields=id,name': [],
+            'fields=id,display': manufacturer,
+            'fields=id,config_context': manufacturer,
+            'fields=id,parent_device': parent,
+            'omit=parent_device': manufacturer,
+            'fields=id,parent_device&omit=parent_device': parent,
+            'brief=1&omit=display': manufacturer + parent,
+            'brief=1&fields=id': [],
+        }
+        for url in (self._get_list_url(), self._get_detail_url(Device.objects.get(name='Device 1'))):
+            for params, expected in cases.items():
+                with self.subTest(url=url, params=params):
+                    with self.capture_querysets(DeviceViewSet) as querysets:
+                        response = self.client.get(f'{url}?{params}', **self.header)
+                    self.assertHttpStatus(response, status.HTTP_200_OK)
+                    self.assert_lookups(querysets, expected)
+        self.assertEqual(DeviceViewSet.queryset._prefetch_related_lookups, ())
+
 
 
 class ModuleTestCase(APIViewTestCases.APIViewTestCase):
 class ModuleTestCase(APIViewTestCases.APIViewTestCase):
     model = Module
     model = Module
@@ -3761,8 +3918,21 @@ class PowerOutletTestCase(Mixins.ComponentTraceMixin, APIViewTestCases.APIViewTe
         ]
         ]
 
 
 
 
-class InterfaceTestCase(Mixins.ComponentTraceMixin, APIViewTestCases.APIViewTestCase):
+class InterfaceTestCase(Mixins.FieldPrefetchMixin, Mixins.ComponentTraceMixin, APIViewTestCases.APIViewTestCase):
     model = Interface
     model = Interface
+    viewset = InterfaceViewSet
+    field_prefetch_cases = (
+        'fields=id,name,device', 'omit=link_peers,link_peers_type', 'fields=id,link_peers&omit=link_peers',
+        'brief=1&omit=l2vpn_termination',
+    )
+    eager_prefetches = (
+        GenericPrefetch('cable__terminations__termination', [Interface.objects.select_related('device', 'cable')]),
+        GenericPrefetch('_path__path_objects', [Interface.objects.select_related('device', 'cable')]),
+        'virtual_circuit_termination',
+        'l2vpn_terminations',
+        'ip_addresses',
+        'fhrp_group_assignments',
+    )
     brief_fields = ['_occupied', 'cable', 'description', 'device', 'display', 'id', 'name', 'url']
     brief_fields = ['_occupied', 'cable', 'description', 'device', 'display', 'id', 'name', 'url']
     bulk_update_data = {
     bulk_update_data = {
         'description': 'New description',
         'description': 'New description',
@@ -3912,6 +4082,53 @@ class InterfaceTestCase(Mixins.ComponentTraceMixin, APIViewTestCases.APIViewTest
             self.assertIn(key, content)
             self.assertIn(key, content)
         self.assertIsNone(content.get('data'))
         self.assertIsNone(content.get('data'))
 
 
+    def _create_field_prefetch_scopes(self):
+        """Return filters for one and for two devices whose interfaces carry every declared prefetch relation."""
+        provider = Provider.objects.create(name='Provider 1', slug='provider-1')
+        provider_network = ProviderNetwork.objects.create(provider=provider, name='Provider Network 1')
+        circuit_type = VirtualCircuitType.objects.create(name='Virtual Circuit Type 1', slug='virtual-circuit-type-1')
+        l2vpn = L2VPN.objects.create(name='L2VPN 1', slug='l2vpn-1', type=L2VPNTypeChoices.TYPE_VXLAN)
+        fhrp_group = FHRPGroup.objects.create(protocol=FHRPGroupProtocolChoices.PROTOCOL_VRRP2, group_id=1)
+        filters = []
+        for i in range(2):
+            device = create_test_device(f'Unit Device {i}')
+            peer = create_test_device(f'Unit Peer {i}')
+            panel = create_test_device(f'Unit Panel {i}')
+            # a: cabled, b: cabled through a patch panel, c: virtual circuit, d: L2VPN, e: IPs and FHRP, f: none
+            interfaces = {
+                name: Interface.objects.create(device=device, name=name, type=InterfaceTypeChoices.TYPE_1GE_FIXED)
+                for name in 'abdef'
+            }
+            interfaces['c'] = Interface.objects.create(device=device, name='c', type=InterfaceTypeChoices.TYPE_VIRTUAL)
+            peers = {
+                name: Interface.objects.create(device=peer, name=name, type=InterfaceTypeChoices.TYPE_1GE_FIXED)
+                for name in 'ab'
+            }
+            peers['c'] = Interface.objects.create(device=peer, name='c', type=InterfaceTypeChoices.TYPE_VIRTUAL)
+            front_port = FrontPort.objects.create(device=panel, name='b', type=PortTypeChoices.TYPE_8P8C)
+            rear_port = RearPort.objects.create(device=panel, name='b', type=PortTypeChoices.TYPE_8P8C)
+            PortMapping.objects.create(device=panel, front_port=front_port, rear_port=rear_port)
+            for a_termination, b_termination in (
+                (interfaces['a'], peers['a']),
+                (interfaces['b'], front_port),
+                (rear_port, peers['b']),
+            ):
+                Cable(a_terminations=[a_termination], b_terminations=[b_termination]).save()
+            virtual_circuit = VirtualCircuit.objects.create(
+                provider_network=provider_network, cid=f'Virtual Circuit {i}', type=circuit_type
+            )
+            for termination in (interfaces['c'], peers['c']):
+                VirtualCircuitTermination.objects.create(
+                    virtual_circuit=virtual_circuit, role=VirtualCircuitTerminationRoleChoices.ROLE_PEER,
+                    interface=termination,
+                )
+            L2VPNTermination.objects.create(l2vpn=l2vpn, assigned_object=interfaces['d'])
+            for j in (1, 2):
+                IPAddress.objects.create(address=f'192.0.2.{10 * i + j}/32', assigned_object=interfaces['e'])
+            FHRPGroupAssignment.objects.create(group=fhrp_group, interface=interfaces['e'], priority=100)
+            filters.append(f'device_id={device.pk}')
+        return [filters[0], '&'.join(filters)]
+
     def test_bulk_delete_child_interfaces(self):
     def test_bulk_delete_child_interfaces(self):
         interface1 = Interface.objects.get(name='Interface 1')
         interface1 = Interface.objects.get(name='Interface 1')
         device = interface1.device
         device = interface1.device
@@ -4356,6 +4573,111 @@ class InterfaceTestCase(Mixins.ComponentTraceMixin, APIViewTestCases.APIViewTest
         response = self.client.post(self._get_list_url(), data, format='json', **self.header)
         response = self.client.post(self._get_list_url(), data, format='json', **self.header)
         self.assertHttpStatus(response, status.HTTP_400_BAD_REQUEST)
         self.assertHttpStatus(response, status.HTTP_400_BAD_REQUEST)
 
 
+    def test_field_prefetches_follow_requested_fields(self):
+        """List and detail responses carry the declared prefetches of exactly the fields they include."""
+        self.add_permissions('dcim.view_interface')
+        declared = InterfaceViewSet.field_prefetches
+        peers = list(declared[('link_peers', 'link_peers_type')])
+        endpoints = list(declared[('connected_endpoints', 'connected_endpoints_type')])
+        path = list(declared[('connected_endpoints_reachable',)])
+        l2vpn = list(declared[('l2vpn_termination',)])
+        ips, fhrp = list(declared[('count_ipaddresses',)]), list(declared[('count_fhrp_groups',)])
+        # fields takes precedence over omit, and both over brief
+        cases = {
+            '': peers + endpoints + path + l2vpn + ips + fhrp,
+            'brief=1': [],
+            'fields=id,link_peers_type': peers,
+            'fields=id,connected_endpoints_type': endpoints,
+            'fields=id,connected_endpoints_reachable': path,
+            'fields=id,l2vpn_termination': l2vpn,
+            'fields=id,count_ipaddresses': ips,
+            'fields=id,count_fhrp_groups': fhrp,
+            'omit=link_peers,link_peers_type': endpoints + path + l2vpn + ips + fhrp,
+            'fields=id,link_peers&omit=link_peers': peers,
+            'brief=1&omit=l2vpn_termination': peers + endpoints + path + ips + fhrp,
+            'brief=1&fields=id': [],
+        }
+        for url in (self._get_list_url(), self._get_detail_url(Interface.objects.get(name='Interface 1'))):
+            for params, expected in cases.items():
+                with self.subTest(url=url, params=params):
+                    with self.capture_querysets(InterfaceViewSet) as querysets:
+                        response = self.client.get(f'{url}?{params}', **self.header)
+                    self.assertHttpStatus(response, status.HTTP_200_OK)
+                    self.assert_lookups(querysets, expected)
+        self.assertEqual(InterfaceViewSet.queryset._prefetch_related_lookups, ())
+
+    def test_field_prefetches_keep_class_lookups(self):
+        """Lookups on the class queryset stay ahead of the declared ones, and a subclass inherits the declarations."""
+        self.add_permissions('dcim.view_interface')
+        url = self._get_list_url()
+        peers = list(InterfaceViewSet.field_prefetches[('link_peers', 'link_peers_type')])
+
+        queryset = InterfaceViewSet.queryset.prefetch_related('wireless_lans')
+        with (
+            patch.object(InterfaceViewSet, 'queryset', queryset),
+            self.capture_querysets(InterfaceViewSet) as querysets,
+        ):
+            response = self.client.get(f'{url}?fields=id,link_peers_type', **self.header)
+        self.assertHttpStatus(response, status.HTTP_200_OK)
+        self.assert_lookups(querysets, [*queryset._prefetch_related_lookups, *peers])
+
+        class PluginInterfaceViewSet(InterfaceViewSet):
+            pass
+
+        with self.capture_querysets(PluginInterfaceViewSet) as querysets:
+            response = PluginInterfaceViewSet.as_view({'get': 'list'})(
+                APIRequestFactory().get(url, {'fields': 'id,link_peers_type'}, **self.header)
+            )
+        self.assertHttpStatus(response, status.HTTP_200_OK)
+        self.assert_lookups(querysets, peers)
+
+    def test_field_prefetches_for_other_reads(self):
+        """Exports and writes carry every declaration, the trace action its path nodes."""
+        self.add_permissions('dcim.view_interface', 'dcim.change_interface', 'extras.view_exporttemplate')
+        declared = [lookup for lookups in InterfaceViewSet.field_prefetches.values() for lookup in lookups]
+        interface = Interface.objects.get(name='Interface 1')
+        export_template = ExportTemplate.objects.create(
+            name='Interfaces', template_code='{% for interface in queryset %}{{ interface.name }},{% endfor %}'
+        )
+        export_template.object_types.set([ObjectType.objects.get_for_model(Interface)])
+
+        # Exports and writes load every declaration, whatever fields the response selects
+        with self.capture_querysets(InterfaceViewSet) as querysets:
+            response = self.client.get(f'{self._get_list_url()}?fields=id&export=Interfaces', **self.header)
+            self.assertHttpStatus(response, status.HTTP_200_OK)
+            response = self.client.patch(
+                f'{self._get_detail_url(interface)}?fields=id', {'description': 'New'}, format='json', **self.header
+            )
+            self.assertHttpStatus(response, status.HTTP_200_OK)
+        self.assert_lookups(querysets, declared)
+
+        # The trace action loads the path nodes as the connected endpoint fields do, whatever fields are selected
+        peer = Interface.objects.create(device=interface.device, name='Peer', type=InterfaceTypeChoices.TYPE_1GE_FIXED)
+        Cable(a_terminations=[interface], b_terminations=[peer]).save()
+        url = reverse('dcim-api:interface-trace', kwargs={'pk': interface.pk})
+
+        def trace(params, queryset, field_prefetches):
+            with (
+                patch.object(InterfaceViewSet, 'queryset', queryset),
+                patch.object(InterfaceViewSet, 'field_prefetches', field_prefetches),
+                CaptureQueriesContext(connection) as queries,
+            ):
+                response = self.client.get(f'{url}?{params}', **self.header)
+            self.assertHttpStatus(response, status.HTTP_200_OK)
+            return response.json(), sum('"users_token"' not in entry['sql'] for entry in queries.captured_queries)
+
+        self.client.get(url, **self.header)
+        queryset = InterfaceViewSet.queryset
+        expected, eager_count = trace('', queryset.prefetch_related(*self.eager_prefetches), {})
+        unprefetched, unprefetched_count = trace('', queryset, {})
+        self.assertEqual(unprefetched, expected)
+        for params in ('', 'fields=id'):
+            with self.subTest(params=params):
+                data, count = trace(params, queryset, InterfaceViewSet.field_prefetches)
+                self.assertEqual(data, expected)
+                self.assertLessEqual(count, eager_count)
+                self.assertLess(count, unprefetched_count)
+
 
 
 class FrontPortTestCase(APIViewTestCases.APIViewTestCase):
 class FrontPortTestCase(APIViewTestCases.APIViewTestCase):
     model = FrontPort
     model = FrontPort

+ 30 - 1
netbox/netbox/api/viewsets/__init__.py

@@ -8,11 +8,12 @@ from django.db.models import ProtectedError, RestrictedError
 from rest_framework import mixins as drf_mixins
 from rest_framework import mixins as drf_mixins
 from rest_framework import status
 from rest_framework import status
 from rest_framework.exceptions import MethodNotAllowed
 from rest_framework.exceptions import MethodNotAllowed
+from rest_framework.permissions import SAFE_METHODS
 from rest_framework.response import Response
 from rest_framework.response import Response
 from rest_framework.viewsets import GenericViewSet
 from rest_framework.viewsets import GenericViewSet
 
 
 from netbox.api.serializers.features import ChangeLogMessageSerializer
 from netbox.api.serializers.features import ChangeLogMessageSerializer
-from utilities.api import get_annotations_for_serializer, get_prefetches_for_serializer
+from utilities.api import get_annotations_for_serializer, get_fields_for_serializer, get_prefetches_for_serializer
 from utilities.exceptions import AbortRequest, PreconditionFailed
 from utilities.exceptions import AbortRequest, PreconditionFailed
 from utilities.query import reapply_model_ordering
 from utilities.query import reapply_model_ordering
 
 
@@ -84,6 +85,8 @@ class BaseViewSet(GenericViewSet):
     Base class for all API ViewSets. This is responsible for the enforcement of object-based permissions.
     Base class for all API ViewSets. This is responsible for the enforcement of object-based permissions.
     """
     """
     brief = False
     brief = False
+    # Prefetches dynamic resolution cannot derive (method fields, model properties), keyed by the fields reading them
+    field_prefetches = {}
 
 
     def initial(self, request, *args, **kwargs):
     def initial(self, request, *args, **kwargs):
         super().initial(request, *args, **kwargs)
         super().initial(request, *args, **kwargs)
@@ -111,6 +114,10 @@ class BaseViewSet(GenericViewSet):
         qs = super().get_queryset()
         qs = super().get_queryset()
         serializer_class = self.get_serializer_class()
         serializer_class = self.get_serializer_class()
 
 
+        # Declared lookups precede the dynamic ones, so a declared Prefetch queryset wins on a shared path
+        if prefetch := self.get_field_prefetches():
+            qs = qs.prefetch_related(*prefetch)
+
         # Dynamically resolve prefetches for included serializer fields and attach them to the queryset
         # Dynamically resolve prefetches for included serializer fields and attach them to the queryset
         if prefetch := get_prefetches_for_serializer(serializer_class, **self.field_kwargs):
         if prefetch := get_prefetches_for_serializer(serializer_class, **self.field_kwargs):
             qs = qs.prefetch_related(*prefetch)
             qs = qs.prefetch_related(*prefetch)
@@ -121,6 +128,23 @@ class BaseViewSet(GenericViewSet):
 
 
         return qs
         return qs
 
 
+    def get_field_prefetches(self, fields=None):
+        """Return the declared prefetches of given fields, by default the response's (all for exports and writes)."""
+        if not self.field_prefetches:
+            return []
+        if fields is None:
+            # Exports and writes render whole objects (templates, event payloads), so they load every declaration
+            if 'export' in self.request.query_params or self.request.method not in SAFE_METHODS:
+                fields = set().union(*self.field_prefetches)
+            else:
+                fields = self.response_fields
+        return [
+            lookup
+            for field_names, lookups in self.field_prefetches.items()
+            if not fields.isdisjoint(field_names)
+            for lookup in lookups
+        ]
+
     def get_serializer(self, *args, **kwargs):
     def get_serializer(self, *args, **kwargs):
         # Pass the fields/omit kwargs (if specified by the request) to the serializer
         # Pass the fields/omit kwargs (if specified by the request) to the serializer
         kwargs.update(**self.field_kwargs)
         kwargs.update(**self.field_kwargs)
@@ -145,6 +169,11 @@ class BaseViewSet(GenericViewSet):
 
 
         return {}
         return {}
 
 
+    @cached_property
+    def response_fields(self):
+        """Return the serializer field names rendered for this request, as selected by field_kwargs."""
+        return frozenset(get_fields_for_serializer(self.get_serializer_class(), **self.field_kwargs))
+
 
 
 class NetBoxReadOnlyModelViewSet(
 class NetBoxReadOnlyModelViewSet(
     ETagMixin,
     ETagMixin,

+ 14 - 15
netbox/utilities/api.py

@@ -30,6 +30,7 @@ logger = logging.getLogger('netbox.utilities.api')
 __all__ = (
 __all__ = (
     'IsSuperuser',
     'IsSuperuser',
     'get_annotations_for_serializer',
     'get_annotations_for_serializer',
+    'get_fields_for_serializer',
     'get_graphql_type_for_model',
     'get_graphql_type_for_model',
     'get_positional_errors',
     'get_positional_errors',
     'get_prefetches_for_serializer',
     'get_prefetches_for_serializer',
@@ -183,19 +184,25 @@ def _get_serializer_fields(serializer: Serializer):
     return [field_name for field_name in fields if field_name not in omit]
     return [field_name for field_name in fields if field_name not in omit]
 
 
 
 
-def get_prefetches_for_serializer(serializer_class, fields=None, omit=None, _serializer_states=None):
+def get_fields_for_serializer(serializer_class, fields=None, omit=None):
     """
     """
-    Compile and return a list of fields which should be prefetched on the queryset for a serializer.
+    Return the names of the fields a serializer renders for the given fields or omit selection.
     """
     """
     if fields is not None and omit is not None:
     if fields is not None and omit is not None:
         raise TypeError("Cannot specify both 'fields' and 'omit' parameters.")
         raise TypeError("Cannot specify both 'fields' and 'omit' parameters.")
 
 
-    model = serializer_class.Meta.model
-
     # If fields are not specified, default to all
     # If fields are not specified, default to all
     fields_to_include = fields or serializer_class.Meta.fields
     fields_to_include = fields or serializer_class.Meta.fields
     fields_to_omit = omit or []
     fields_to_omit = omit or []
-    effective_fields = tuple(name for name in fields_to_include if name not in fields_to_omit)
+    return tuple(name for name in fields_to_include if name not in fields_to_omit)
+
+
+def get_prefetches_for_serializer(serializer_class, fields=None, omit=None, _serializer_states=None):
+    """
+    Compile and return a list of fields which should be prefetched on the queryset for a serializer.
+    """
+    effective_fields = get_fields_for_serializer(serializer_class, fields, omit)
+    model = serializer_class.Meta.model
 
 
     # Break reference cycles on the current path. The field set is in the key because re-entry at a
     # Break reference cycles on the current path. The field set is in the key because re-entry at a
     # narrower depth is finite, and the states are copied per frame to keep sibling fields independent.
     # narrower depth is finite, and the states are copied per frame to keep sibling fields independent.
@@ -239,20 +246,12 @@ def get_annotations_for_serializer(serializer_class, fields=None, omit=None):
     """
     """
     Return a mapping of field names to annotations to be applied to the queryset for a serializer.
     Return a mapping of field names to annotations to be applied to the queryset for a serializer.
     """
     """
-    if fields is not None and omit is not None:
-        raise TypeError("Cannot specify both 'fields' and 'omit' parameters.")
-
+    effective_fields = get_fields_for_serializer(serializer_class, fields, omit)
     model = serializer_class.Meta.model
     model = serializer_class.Meta.model
 
 
-    # If fields are not specified, default to all
-    fields_to_include = fields or serializer_class.Meta.fields
-    fields_to_omit = omit or []
-
     annotations = {}
     annotations = {}
     for field_name, field in serializer_class._declared_fields.items():
     for field_name, field in serializer_class._declared_fields.items():
-        if field_name in fields_to_omit:
-            continue
-        if field_name in fields_to_include and type(field) is RelatedObjectCountField:
+        if field_name in effective_fields and type(field) is RelatedObjectCountField:
             related_field = getattr(model, field.relation).field
             related_field = getattr(model, field.relation).field
             annotations[field_name] = count_related(related_field.model, related_field.name)
             annotations[field_name] = count_related(related_field.model, related_field.name)