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

fix(api): Discover nested prefetches for SerializedPKRelatedField

Teach _get_nested_serializer() to unwrap ManyRelatedField and resolve the
serializer configured on SerializedPKRelatedField. Preserve the field's
nested value so prefetch discovery follows brief_fields for nested
representations and the full field set otherwise.

Guard the recursion against reference cycles by tracking the serializer
states already being resolved on the current path. Because nested defaults
to False the target expands its full field set, so a self-referential or
mutually referential declaration would otherwise recurse without a
termination condition. The key pairs the serializer class with its
effective field set, since re-entering a serializer at a narrower depth is
finite rather than cyclic, and the states are copied per frame so a sibling
field can still expand what another one stops at.

Populate interface VDC relationships and VRF route target assignments in
the API test fixtures, update the two query count baselines, and add
coverage for direct, many-valued and non-nested serialized related fields
alongside self-referential, mutually referential, brief-depth and
sibling-reuse cases.

Fixes #22988
Martin Hauser 20 часов назад
Родитель
Сommit
2b54582cd7

+ 1 - 1
netbox/dcim/tests/query_counts.json

@@ -23,7 +23,7 @@
   "frontport:api_list_objects": 14,
   "frontport:api_list_objects": 14,
   "frontport:list_objects_with_permission": 22,
   "frontport:list_objects_with_permission": 22,
   "frontporttemplate:api_list_objects": 12,
   "frontporttemplate:api_list_objects": 12,
-  "interface:api_list_objects": 22,
+  "interface:api_list_objects": 23,
   "interface:list_objects_with_permission": 18,
   "interface:list_objects_with_permission": 18,
   "interfaceconnection:list_objects_with_permission": 41,
   "interfaceconnection:list_objects_with_permission": 41,
   "interfacetemplate:api_list_objects": 11,
   "interfacetemplate:api_list_objects": 11,

+ 2 - 0
netbox/dcim/tests/test_api.py

@@ -2711,6 +2711,8 @@ class InterfaceTestCase(Mixins.ComponentTraceMixin, APIViewTestCases.APIViewTest
             VirtualDeviceContext(name='VDC 2', identifier=2, device=device)
             VirtualDeviceContext(name='VDC 2', identifier=2, device=device)
         )
         )
         VirtualDeviceContext.objects.bulk_create(vdcs)
         VirtualDeviceContext.objects.bulk_create(vdcs)
+        for interface in interfaces:
+            interface.vdcs.set(vdcs)
 
 
         vlans = (
         vlans = (
             VLAN(name='VLAN 1', vid=1),
             VLAN(name='VLAN 1', vid=1),

+ 1 - 1
netbox/ipam/tests/query_counts.json

@@ -32,6 +32,6 @@
   "vlantranslationpolicy:list_objects_with_permission": 17,
   "vlantranslationpolicy:list_objects_with_permission": 17,
   "vlantranslationrule:api_list_objects": 12,
   "vlantranslationrule:api_list_objects": 12,
   "vlantranslationrule:list_objects_with_permission": 18,
   "vlantranslationrule:list_objects_with_permission": 18,
-  "vrf:api_list_objects": 14,
+  "vrf:api_list_objects": 20,
   "vrf:list_objects_with_permission": 17
   "vrf:list_objects_with_permission": 17
 }
 }

+ 14 - 0
netbox/ipam/tests/test_api.py

@@ -224,6 +224,15 @@ class VRFTestCase(APIViewTestCases.APIViewTestCase):
     @classmethod
     @classmethod
     def setUpTestData(cls):
     def setUpTestData(cls):
 
 
+        tenant = Tenant.objects.create(name='Tenant 1', slug='tenant-1')
+
+        route_targets = (
+            RouteTarget(name='65000:1001', tenant=tenant),
+            RouteTarget(name='65000:1002', tenant=tenant),
+            RouteTarget(name='65000:1003', tenant=tenant),
+        )
+        RouteTarget.objects.bulk_create(route_targets)
+
         vrfs = (
         vrfs = (
             VRF(name='VRF 1', rd='65000:1'),
             VRF(name='VRF 1', rd='65000:1'),
             VRF(name='VRF 2', rd='65000:2'),
             VRF(name='VRF 2', rd='65000:2'),
@@ -231,6 +240,11 @@ class VRFTestCase(APIViewTestCases.APIViewTestCase):
         )
         )
         VRF.objects.bulk_create(vrfs)
         VRF.objects.bulk_create(vrfs)
 
 
+        # Assigned so the query count baseline covers the non-nested route target expansion.
+        for vrf in vrfs:
+            vrf.import_targets.set(route_targets)
+            vrf.export_targets.set(route_targets)
+
 
 
 class RouteTargetTestCase(APIViewTestCases.APIViewTestCase):
 class RouteTargetTestCase(APIViewTestCases.APIViewTestCase):
     model = RouteTarget
     model = RouteTarget

+ 23 - 6
netbox/utilities/api.py

@@ -13,12 +13,13 @@ from django.urls import reverse
 from django.utils.module_loading import import_string
 from django.utils.module_loading import import_string
 from django.utils.translation import gettext_lazy as _
 from django.utils.translation import gettext_lazy as _
 from rest_framework.permissions import BasePermission
 from rest_framework.permissions import BasePermission
+from rest_framework.relations import ManyRelatedField
 from rest_framework.serializers import ListSerializer, Serializer
 from rest_framework.serializers import ListSerializer, Serializer
 from rest_framework.views import get_view_name as drf_get_view_name
 from rest_framework.views import get_view_name as drf_get_view_name
 
 
 from extras.constants import HTTP_CONTENT_TYPE_JSON
 from extras.constants import HTTP_CONTENT_TYPE_JSON
 from netbox.api.exceptions import GraphQLTypeNotFound, SerializerNotFound
 from netbox.api.exceptions import GraphQLTypeNotFound, SerializerNotFound
-from netbox.api.fields import RelatedObjectCountField
+from netbox.api.fields import RelatedObjectCountField, SerializedPKRelatedField
 from netbox.registry import registry
 from netbox.registry import registry
 
 
 from .query import count_related, dict_to_filter_params
 from .query import count_related, dict_to_filter_params
@@ -134,6 +135,13 @@ def _get_nested_serializer(serializer_field):
     if isinstance(serializer_field, ListSerializer):
     if isinstance(serializer_field, ListSerializer):
         serializer_field = serializer_field.child
         serializer_field = serializer_field.child
 
 
+    # DRF wraps a many-valued related field, keeping the original field on child_relation
+    if isinstance(serializer_field, ManyRelatedField):
+        serializer_field = serializer_field.child_relation
+
+    if isinstance(serializer_field, SerializedPKRelatedField):
+        return serializer_field.serializer(nested=serializer_field.nested)
+
     if isinstance(serializer_field, Serializer) and hasattr(serializer_field, 'nested'):
     if isinstance(serializer_field, Serializer) and hasattr(serializer_field, 'nested'):
         return serializer_field
         return serializer_field
 
 
@@ -151,7 +159,7 @@ 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):
+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.
     Compile and return a list of fields which should be prefetched on the queryset for a serializer.
     """
     """
@@ -163,11 +171,18 @@ def get_prefetches_for_serializer(serializer_class, fields=None, omit=None):
     # 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)
+
+    # 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.
+    serializer_states = set(_serializer_states or ())
+    serializer_state = (serializer_class, effective_fields)
+    if serializer_state in serializer_states:
+        return []
+    serializer_states.add(serializer_state)
 
 
     prefetch_fields = []
     prefetch_fields = []
-    for field_name in fields_to_include:
-        if field_name in fields_to_omit:
-            continue
+    for field_name in effective_fields:
         serializer_field = serializer_class._declared_fields.get(field_name)
         serializer_field = serializer_class._declared_fields.get(field_name)
 
 
         # Determine the name of the model field referenced by the serializer field
         # Determine the name of the model field referenced by the serializer field
@@ -188,7 +203,9 @@ def get_prefetches_for_serializer(serializer_class, fields=None, omit=None):
         # constraints set on that serializer field instance.
         # constraints set on that serializer field instance.
         if nested_serializer := _get_nested_serializer(serializer_field):
         if nested_serializer := _get_nested_serializer(serializer_field):
             subfields = _get_serializer_fields(nested_serializer)
             subfields = _get_serializer_fields(nested_serializer)
-            for subfield in get_prefetches_for_serializer(type(nested_serializer), fields=subfields):
+            for subfield in get_prefetches_for_serializer(
+                type(nested_serializer), fields=subfields, _serializer_states=serializer_states
+            ):
                 prefetch_fields.append(f'{field.name}__{subfield}')
                 prefetch_fields.append(f'{field.name}__{subfield}')
 
 
     return prefetch_fields
     return prefetch_fields

+ 173 - 0
netbox/utilities/tests/test_api.py

@@ -11,6 +11,7 @@ from extras.choices import CustomFieldTypeChoices
 from extras.models import CustomField
 from extras.models import CustomField
 from ipam.api.serializers import VLANSerializer
 from ipam.api.serializers import VLANSerializer
 from ipam.models import VLAN
 from ipam.models import VLAN
+from netbox.api.fields import SerializedPKRelatedField
 from netbox.api.serializers import BaseModelSerializer
 from netbox.api.serializers import BaseModelSerializer
 from netbox.config import get_config
 from netbox.config import get_config
 from netbox.plugins import register_serializer_resolver
 from netbox.plugins import register_serializer_resolver
@@ -581,6 +582,178 @@ class GetPrefetchesForSerializerTestCase(TestCase):
             ['region', 'region__parent'],
             ['region', 'region__parent'],
         )
         )
 
 
+    def test_serialized_pk_related_field(self):
+        class RegionSerializer(BaseModelSerializer):
+            class Meta:
+                model = Region
+                fields = ('id', 'name', 'parent', 'sites')
+                brief_fields = ('id', 'parent')
+
+        class SiteSerializer(BaseModelSerializer):
+            region = SerializedPKRelatedField(
+                queryset=Region.objects.all(),
+                serializer=RegionSerializer,
+                nested=True,
+            )
+
+            class Meta:
+                model = Site
+                fields = ('id', 'region')
+
+        self.assertListEqual(
+            get_prefetches_for_serializer(SiteSerializer),
+            ['region', 'region__parent'],
+        )
+
+    def test_many_serialized_pk_related_field(self):
+        class SiteSerializer(BaseModelSerializer):
+            class Meta:
+                model = Site
+                fields = ('id', 'name', 'region', 'group')
+                brief_fields = ('id', 'region')
+
+        class RegionSerializer(BaseModelSerializer):
+            sites = SerializedPKRelatedField(
+                queryset=Site.objects.all(),
+                serializer=SiteSerializer,
+                nested=True,
+                many=True,
+            )
+
+            class Meta:
+                model = Region
+                fields = ('id', 'sites')
+
+        self.assertListEqual(
+            get_prefetches_for_serializer(RegionSerializer),
+            ['sites', 'sites__region'],
+        )
+
+        self.assertListEqual(
+            get_prefetches_for_serializer(RegionSerializer, fields=('id',)),
+            [],
+        )
+
+        self.assertListEqual(
+            get_prefetches_for_serializer(RegionSerializer, omit=('sites',)),
+            [],
+        )
+
+    def test_many_serialized_pk_related_field_not_nested(self):
+        class SiteSerializer(BaseModelSerializer):
+            class Meta:
+                model = Site
+                fields = ('id', 'name', 'region', 'group')
+                brief_fields = ('id', 'region')
+
+        class RegionSerializer(BaseModelSerializer):
+            sites = SerializedPKRelatedField(
+                queryset=Site.objects.all(),
+                serializer=SiteSerializer,
+                nested=False,
+                many=True,
+            )
+
+            class Meta:
+                model = Region
+                fields = ('id', 'sites')
+
+        self.assertListEqual(
+            get_prefetches_for_serializer(RegionSerializer),
+            ['sites', 'sites__region', 'sites__group'],
+        )
+
+    def test_self_referential_serialized_pk_related_field(self):
+        class RegionSerializer(BaseModelSerializer):
+            class Meta:
+                model = Region
+                fields = ('id', 'parent', 'children')
+
+        # The field can only name its own serializer once the class exists.
+        RegionSerializer._declared_fields['children'] = SerializedPKRelatedField(
+            queryset=Region.objects.all(),
+            serializer=RegionSerializer,
+            many=True,
+        )
+
+        self.assertListEqual(
+            get_prefetches_for_serializer(RegionSerializer),
+            ['parent', 'children'],
+        )
+
+    def test_self_referential_serialized_pk_related_field_with_brief_fields(self):
+        class RegionSerializer(BaseModelSerializer):
+            class Meta:
+                model = Region
+                fields = ('id', 'sites', 'children')
+                brief_fields = ('id', 'sites')
+
+        RegionSerializer._declared_fields['children'] = SerializedPKRelatedField(
+            queryset=Region.objects.all(),
+            serializer=RegionSerializer,
+            nested=True,
+            many=True,
+        )
+
+        # Re-entering the serializer at brief depth is not a cycle, so brief_fields must expand.
+        self.assertListEqual(
+            get_prefetches_for_serializer(RegionSerializer),
+            ['sites', 'children', 'children__sites'],
+        )
+
+    def test_mutually_referential_serialized_pk_related_fields(self):
+        class RegionSerializer(BaseModelSerializer):
+            class Meta:
+                model = Region
+                fields = ('id', 'sites')
+
+        class SiteSerializer(BaseModelSerializer):
+            region = SerializedPKRelatedField(
+                queryset=Region.objects.all(),
+                serializer=RegionSerializer,
+            )
+
+            class Meta:
+                model = Site
+                fields = ('id', 'region')
+
+        RegionSerializer._declared_fields['sites'] = SerializedPKRelatedField(
+            queryset=Site.objects.all(),
+            serializer=SiteSerializer,
+            many=True,
+        )
+
+        self.assertListEqual(
+            get_prefetches_for_serializer(RegionSerializer),
+            ['sites', 'sites__region'],
+        )
+
+    def test_serializer_class_reused_on_sibling_fields(self):
+        class TargetRegionSerializer(BaseModelSerializer):
+            class Meta:
+                model = Region
+                fields = ('id', 'sites')
+
+        class RegionSerializer(BaseModelSerializer):
+            parent = SerializedPKRelatedField(
+                queryset=Region.objects.all(),
+                serializer=TargetRegionSerializer,
+            )
+            children = SerializedPKRelatedField(
+                queryset=Region.objects.all(),
+                serializer=TargetRegionSerializer,
+                many=True,
+            )
+
+            class Meta:
+                model = Region
+                fields = ('id', 'parent', 'children')
+
+        self.assertListEqual(
+            get_prefetches_for_serializer(RegionSerializer),
+            ['parent', 'parent__sites', 'children', 'children__sites'],
+        )
+
 
 
 class _ResolvedSerializerA(Serializer):
 class _ResolvedSerializerA(Serializer):
     pass
     pass