perf: batch export graph database queries

This commit is contained in:
2026-08-05 12:55:37 +02:00
parent 80ef1be239
commit f5ad77b430
10 changed files with 331 additions and 128 deletions
+5
View File
@@ -42,6 +42,7 @@ PLUGINS_CONFIG = {
"netbox_export": { "netbox_export": {
"max_objects": 50000, "max_objects": 50000,
"max_archive_size_mb": 250, "max_archive_size_mb": 250,
"query_batch_size": 500,
# Auf beiden Instanzen identisch setzen, um Archive zu signieren. # Auf beiden Instanzen identisch setzen, um Archive zu signieren.
"archive_signing_key": "eine-lange-zufaellige-geheime-zeichenfolge", "archive_signing_key": "eine-lange-zufaellige-geheime-zeichenfolge",
}, },
@@ -99,6 +100,10 @@ Import werden vorhandene Objekte über ihre eindeutigen Fachschlüssel erkannt.
werden. werden.
- Große Exporte werden synchron verarbeitet. `max_objects` begrenzt Laufzeit und - Große Exporte werden synchron verarbeitet. `max_objects` begrenzt Laufzeit und
Speicherverbrauch. Speicherverbrauch.
- `query_batch_size` steuert die Größe gebündelter Datenbankabfragen. Der
Standardwert `500` ist für typische PostgreSQL-Installationen geeignet;
Werte zwischen `250` und `1000` erlauben eine Anpassung an Arbeitsspeicher und
Datenbankleistung.
## Tests ## Tests
+2 -1
View File
@@ -7,7 +7,7 @@ class NetBoxExportConfig(PluginConfig):
name = "netbox_export" name = "netbox_export"
verbose_name = "NetBox-Export" verbose_name = "NetBox-Export"
description = "Portable ZIP export and import for tenants and locations" description = "Portable ZIP export and import for tenants and locations"
version = "0.2.0" version = "0.3.0"
author = "NetBox Export contributors" author = "NetBox Export contributors"
base_url = "netbox-export" base_url = "netbox-export"
min_version = "4.6.0" min_version = "4.6.0"
@@ -16,6 +16,7 @@ class NetBoxExportConfig(PluginConfig):
default_settings: ClassVar[dict] = { default_settings: ClassVar[dict] = {
"max_objects": 50000, "max_objects": 50000,
"max_archive_size_mb": 250, "max_archive_size_mb": 250,
"query_batch_size": 500,
"archive_signing_key": "", "archive_signing_key": "",
} }
+26 -2
View File
@@ -8,6 +8,7 @@ import posixpath
import zipfile import zipfile
from collections.abc import Iterable from collections.abc import Iterable
from dataclasses import dataclass from dataclasses import dataclass
from pathlib import PurePosixPath
from typing import BinaryIO from typing import BinaryIO
from django.core.serializers.json import DjangoJSONEncoder from django.core.serializers.json import DjangoJSONEncoder
@@ -18,6 +19,19 @@ FORMAT_NAME = "netbox-export"
FORMAT_VERSION = 1 FORMAT_VERSION = 1
MANIFEST_NAME = "manifest.json" MANIFEST_NAME = "manifest.json"
OBJECTS_NAME = "objects.ndjson" OBJECTS_NAME = "objects.ndjson"
PRECOMPRESSED_SUFFIXES = {
".7z",
".avi",
".gz",
".jpeg",
".jpg",
".mp3",
".mp4",
".pdf",
".png",
".webp",
".zip",
}
def _json_bytes(value) -> bytes: def _json_bytes(value) -> bytes:
@@ -59,11 +73,21 @@ def build_archive(manifest: dict, records: Iterable[dict], assets: dict[str, byt
manifest["signature"] = _signature(manifest, signing_key) manifest["signature"] = _signature(manifest, signing_key)
output = io.BytesIO() output = io.BytesIO()
with zipfile.ZipFile(output, "w", compression=zipfile.ZIP_DEFLATED, compresslevel=6) as archive: with zipfile.ZipFile(output, "w", compression=zipfile.ZIP_DEFLATED, compresslevel=1) as archive:
archive.writestr(MANIFEST_NAME, _json_bytes(manifest)) archive.writestr(MANIFEST_NAME, _json_bytes(manifest))
archive.writestr(OBJECTS_NAME, object_data) archive.writestr(OBJECTS_NAME, object_data)
for path, content in sorted(assets.items()): for path, content in sorted(assets.items()):
archive.writestr(_safe_member_name(path), content) compression = (
zipfile.ZIP_STORED
if PurePosixPath(path).suffix.lower() in PRECOMPRESSED_SUFFIXES
else zipfile.ZIP_DEFLATED
)
archive.writestr(
_safe_member_name(path),
content,
compress_type=compression,
compresslevel=None if compression == zipfile.ZIP_STORED else 1,
)
return output.getvalue() return output.getvalue()
+32 -32
View File
@@ -87,16 +87,23 @@ def external_identity(obj) -> dict:
return {"model": label, "lookup": {"pk": encode_scalar(obj.pk)}} return {"model": label, "lookup": {"pk": encode_scalar(obj.pk)}}
def _reference_spec(obj, exported_keys: set[str]): def _reference_for_pk(model, pk, exported_keys: set[str]):
if object_key(obj) in exported_keys: key = f"{model._meta.label_lower}:{pk}"
return {"ref": object_key(obj)} if key in exported_keys:
return {"external": external_identity(obj)} return {"ref": key}
try:
target = model._default_manager.get(pk=pk)
except model.DoesNotExist:
return None
return {"external": external_identity(target)}
def _encode_custom_field_data(obj, value: dict, exported_keys: set[str]): def _encode_custom_field_data(obj, value: dict, exported_keys: set[str], custom_fields=None):
from extras.models import CustomField if custom_fields is None:
from extras.models import CustomField
definitions = {field.name: field for field in CustomField.objects.get_for_model(type(obj))} custom_fields = CustomField.objects.get_for_model(type(obj))
definitions = {field.name: field for field in custom_fields}
encoded = {} encoded = {}
for name, raw_value in value.items(): for name, raw_value in value.items():
custom_field = definitions.get(name) custom_field = definitions.get(name)
@@ -110,21 +117,17 @@ def _encode_custom_field_data(obj, value: dict, exported_keys: set[str]):
continue continue
target_model = custom_field.related_object_type.model_class() target_model = custom_field.related_object_type.model_class()
if custom_field.type == "object": if custom_field.type == "object":
try: reference = _reference_for_pk(target_model, raw_value, exported_keys)
target = target_model._default_manager.get(pk=raw_value) encoded[name] = {"$type": "object_ref", "value": reference} if reference else None
except target_model.DoesNotExist:
encoded[name] = None
else:
encoded[name] = {"$type": "object_ref", "value": _reference_spec(target, exported_keys)}
else: else:
targets = {str(item.pk): item for item in target_model._default_manager.filter(pk__in=raw_value)} references = [
reference
for pk in raw_value
if (reference := _reference_for_pk(target_model, pk, exported_keys)) is not None
]
encoded[name] = { encoded[name] = {
"$type": "multiobject_ref", "$type": "multiobject_ref",
"value": [ "value": references,
_reference_spec(targets[str(pk)], exported_keys)
for pk in raw_value
if str(pk) in targets
],
} }
return encoded return encoded
@@ -138,23 +141,20 @@ def _encode_custom_field_default(custom_field, value, exported_keys: set[str]):
return encode_scalar(value) return encode_scalar(value)
target_model = custom_field.related_object_type.model_class() target_model = custom_field.related_object_type.model_class()
if custom_field.type == "object": if custom_field.type == "object":
try: reference = _reference_for_pk(target_model, value, exported_keys)
target = target_model._default_manager.get(pk=value) return {"$type": "object_ref", "value": reference} if reference else None
except target_model.DoesNotExist: references = [
return None reference
return {"$type": "object_ref", "value": _reference_spec(target, exported_keys)} for pk in value
targets = {str(item.pk): item for item in target_model._default_manager.filter(pk__in=value)} if (reference := _reference_for_pk(target_model, pk, exported_keys)) is not None
]
return { return {
"$type": "multiobject_ref", "$type": "multiobject_ref",
"value": [ "value": references,
_reference_spec(targets[str(pk)], exported_keys)
for pk in value
if str(pk) in targets
],
} }
def serialize_object(obj, exported_keys: set[str], assets: dict[str, bytes]) -> dict: def serialize_object(obj, exported_keys: set[str], assets: dict[str, bytes], custom_fields=None) -> dict:
record_id = object_key(obj) record_id = object_key(obj)
gfk_fields = generic_foreign_keys(type(obj)) gfk_fields = generic_foreign_keys(type(obj))
gfk_storage = {name for field in gfk_fields for name in (field.ct_field, field.fk_field)} gfk_storage = {name for field in gfk_fields for name in (field.ct_field, field.fk_field)}
@@ -191,7 +191,7 @@ def serialize_object(obj, exported_keys: set[str], assets: dict[str, bytes]) ->
continue continue
value = field.value_from_object(obj) value = field.value_from_object(obj)
if field.name == "custom_field_data" and isinstance(value, dict): if field.name == "custom_field_data" and isinstance(value, dict):
scalars[field.name] = _encode_custom_field_data(obj, value, exported_keys) scalars[field.name] = _encode_custom_field_data(obj, value, exported_keys, custom_fields)
elif obj._meta.label_lower == "extras.customfield" and field.name == "default": elif obj._meta.label_lower == "extras.customfield" and field.name == "default":
scalars[field.name] = _encode_custom_field_default(obj, value, exported_keys) scalars[field.name] = _encode_custom_field_default(obj, value, exported_keys)
else: else:
+18 -5
View File
@@ -12,21 +12,34 @@ from .codec import serialize_object
from .graph import ObjectGraph, seed_scope from .graph import ObjectGraph, seed_scope
def export_scope(scope_type: str, scope_id: int, *, max_objects: int, signing_key: str = ""): def export_scope(
scope_type: str,
scope_id: int,
*,
max_objects: int,
query_batch_size: int = 500,
signing_key: str = "",
):
root, seeds = seed_scope(scope_type, scope_id) root, seeds = seed_scope(scope_type, scope_id)
graph = ObjectGraph(max_objects=max_objects).collect(seeds) graph = ObjectGraph(max_objects=max_objects, query_batch_size=query_batch_size).collect(seeds)
objects = graph.objects objects = graph.objects
exported_keys = set(objects)
assets = {} assets = {}
records = [ records = [
serialize_object(obj, set(objects), assets) serialize_object(
for _, obj in sorted(objects.items()) obj,
exported_keys,
assets,
custom_fields=graph.custom_fields_for_model(type(obj)),
)
for obj in graph.iter_loaded_objects()
] ]
counts = Counter(record["model"] for record in records) counts = Counter(record["model"] for record in records)
manifest = { manifest = {
"created_at": datetime.now(UTC).isoformat(), "created_at": datetime.now(UTC).isoformat(),
"source_instance": str(InstanceIdentity.local_id()), "source_instance": str(InstanceIdentity.local_id()),
"source_netbox_version": getattr(getattr(settings, "RELEASE", None), "version", "4.6"), "source_netbox_version": getattr(getattr(settings, "RELEASE", None), "version", "4.6"),
"plugin_version": "0.2.0", "plugin_version": "0.3.0",
"scope": { "scope": {
"type": scope_type, "type": scope_type,
"source_pk": str(scope_id), "source_pk": str(scope_id),
+178 -87
View File
@@ -1,6 +1,6 @@
from __future__ import annotations from __future__ import annotations
from collections import deque from collections import defaultdict
from django.apps import apps from django.apps import apps
from django.contrib.contenttypes.fields import GenericForeignKey from django.contrib.contenttypes.fields import GenericForeignKey
@@ -55,6 +55,12 @@ def _descendants(obj):
return [obj] return [obj]
def _chunks(values, size):
values = list(values)
for offset in range(0, len(values), size):
yield values[offset : offset + size]
def seed_scope(scope_type: str, scope_id: int): def seed_scope(scope_type: str, scope_id: int):
from dcim.models import Location, Region, Site from dcim.models import Location, Region, Site
from tenancy.models import Tenant, TenantGroup from tenancy.models import Tenant, TenantGroup
@@ -80,19 +86,45 @@ def seed_scope(scope_type: str, scope_id: int):
class ObjectGraph: class ObjectGraph:
"""Collect scoped objects first and their forward dependencies second.""" """Collect scoped objects and dependencies using batched relation queries."""
def __init__(self, max_objects: int): def __init__(self, max_objects: int, query_batch_size: int = 500):
self.max_objects = max_objects self.max_objects = max_objects
self.query_batch_size = max(1, query_batch_size)
self.members: dict[str, models.Model] = {} self.members: dict[str, models.Model] = {}
self.dependencies: dict[str, models.Model] = {} self.dependencies: dict[str, models.Model] = {}
self._custom_field_cache = {}
self._models = exportable_models()
self._reverse_fields, self._generic_fields, self._peer_fields = self._build_relation_indexes()
@property @property
def objects(self) -> dict[str, models.Model]: def objects(self) -> dict[str, models.Model]:
return {**self.members, **self.dependencies} return {**self.members, **self.dependencies}
def _build_relation_indexes(self):
reverse_fields = defaultdict(lambda: defaultdict(list))
generic_fields = defaultdict(list)
peer_fields = defaultdict(list)
for candidate_model in self._models:
for relation_field in candidate_model._meta.concrete_fields:
if not isinstance(relation_field, (models.ForeignKey, models.OneToOneField)):
continue
related_model = relation_field.remote_field.model
if (
relation_field.remote_field.on_delete in (models.CASCADE, models.PROTECT)
or relation_field.name in SCOPE_LINK_FIELDS
):
reverse_fields[related_model][candidate_model].append(relation_field)
if related_model._meta.label_lower in PEER_CONTAINER_MODELS:
peer_fields[candidate_model].append(relation_field)
for generic_field in candidate_model._meta.private_fields:
if isinstance(generic_field, GenericForeignKey):
generic_fields[candidate_model].append(generic_field)
return reverse_fields, generic_fields, peer_fields
def _check_limit(self): def _check_limit(self):
if len(self.objects) > self.max_objects: object_count = len(self.members) + len(self.dependencies)
if object_count > self.max_objects:
raise GraphLimitError( raise GraphLimitError(
f"Der Export würde mehr als {self.max_objects} Objekte enthalten. " f"Der Export würde mehr als {self.max_objects} Objekte enthalten. "
"Bitte den Bereich verkleinern oder max_objects erhöhen." "Bitte den Bereich verkleinern oder max_objects erhöhen."
@@ -115,92 +147,123 @@ class ObjectGraph:
self._check_limit() self._check_limit()
return True return True
def collect(self, seeds): def _group_by_model(self, objects):
queue = deque() grouped = defaultdict(list)
for obj in objects:
grouped[type(obj)].append(obj)
return grouped
def _expand_members(self, seeds):
pending = {}
for obj in seeds: for obj in seeds:
if self.add_member(obj): if self.add_member(obj):
queue.append(obj) pending[object_key(obj)] = obj
models_to_scan = exportable_models() while pending:
while True: next_pending = {}
while queue: for parent_model, parents in self._group_by_model(pending.values()).items():
parent = queue.popleft()
parent_model = type(parent)
parent_ct = ContentType.objects.get_for_model(parent_model) parent_ct = ContentType.objects.get_for_model(parent_model)
for candidate_model in models_to_scan: direct_candidates = self._reverse_fields.get(parent_model, {})
query = models.Q() peer_fields = self._peer_fields.get(parent_model, [])
for field in candidate_model._meta.concrete_fields: for parent_batch in _chunks(parents, self.query_batch_size):
if not isinstance(field, (models.ForeignKey, models.OneToOneField)): parent_ids = [obj.pk for obj in parent_batch]
continue
if field.remote_field.model is not parent_model:
continue
if (
field.remote_field.on_delete not in (models.CASCADE, models.PROTECT)
and field.name not in SCOPE_LINK_FIELDS
):
continue
query |= models.Q(**{field.attname: parent.pk})
for field in candidate_model._meta.private_fields:
if isinstance(field, GenericForeignKey):
query |= models.Q(**{field.ct_field: parent_ct, field.fk_field: parent.pk})
if not query:
continue
for child in candidate_model.objects.filter(query).distinct().iterator():
if self.add_member(child):
queue.append(child)
promoted = False if peer_fields:
for obj in list(self.members.values()): peer_names = [field.name for field in peer_fields]
for field in obj._meta.concrete_fields: queryset = parent_model._default_manager.filter(pk__in=parent_ids).select_related(*peer_names)
if not isinstance(field, (models.ForeignKey, models.OneToOneField)): for parent in queryset.iterator(chunk_size=self.query_batch_size):
continue for relation_field in peer_fields:
related = getattr(obj, field.name, None) related = getattr(parent, relation_field.name, None)
if related is None or related._meta.label_lower not in PEER_CONTAINER_MODELS: if related is not None and self.add_member(related):
continue next_pending[object_key(related)] = related
if self.add_member(related):
queue.append(related)
promoted = True
if not promoted:
break
dependency_queue = deque(self.members.values()) candidate_queries = defaultdict(models.Q)
scanned = set() for candidate_model, relation_fields in direct_candidates.items():
while dependency_queue: for relation_field in relation_fields:
obj = dependency_queue.popleft() candidate_queries[candidate_model] |= models.Q(**{f"{relation_field.attname}__in": parent_ids})
key = object_key(obj) for candidate_model, generic_fields in self._generic_fields.items():
if key in scanned: for generic_field in generic_fields:
continue candidate_queries[candidate_model] |= models.Q(
scanned.add(key) **{
related_objects = [] generic_field.ct_field: parent_ct,
for field in obj._meta.concrete_fields: f"{generic_field.fk_field}__in": parent_ids,
if isinstance(field, (models.ForeignKey, models.OneToOneField)): }
related = getattr(obj, field.name, None) )
if related is not None:
related_objects.append(related) for candidate_model, query in candidate_queries.items():
for field in obj._meta.private_fields: queryset = candidate_model._default_manager.filter(query).distinct()
if isinstance(field, GenericForeignKey): for child in queryset.iterator(chunk_size=self.query_batch_size):
related = getattr(obj, field.name, None) if self.add_member(child):
if related is not None: next_pending[object_key(child)] = child
related_objects.append(related) pending = next_pending
for field in obj._meta.many_to_many:
try: def custom_fields_for_model(self, model):
related_objects.extend(getattr(obj, field.name).all()) if model not in self._custom_field_cache:
except (AttributeError, TypeError): if hasattr(model, "custom_field_data"):
pass
if hasattr(obj, "custom_field_data"):
from extras.models import CustomField from extras.models import CustomField
custom_fields = list(CustomField.objects.get_for_model(type(obj))) self._custom_field_cache[model] = list(CustomField.objects.get_for_model(model))
related_objects.extend(custom_fields) else:
for custom_field in custom_fields: self._custom_field_cache[model] = []
if custom_field.type not in ("object", "multiobject") or not custom_field.related_object_type: return self._custom_field_cache[model]
continue
raw_value = obj.custom_field_data.get(custom_field.name) def _load_objects(self, model, objects):
if raw_value in (None, "", []): foreign_keys = [
continue field.name
target_model = custom_field.related_object_type.model_class() for field in model._meta.concrete_fields
target_ids = raw_value if custom_field.type == "multiobject" else [raw_value] if isinstance(field, (models.ForeignKey, models.OneToOneField))
related_objects.extend(target_model._default_manager.filter(pk__in=target_ids)) ]
prefetch_fields = [field.name for field in model._meta.many_to_many]
prefetch_fields.extend(
field.name for field in model._meta.private_fields if isinstance(field, GenericForeignKey)
)
for object_batch in _chunks(objects, self.query_batch_size):
object_ids = [obj.pk for obj in object_batch]
queryset = model._default_manager.filter(pk__in=object_ids).select_related(*foreign_keys).order_by("pk")
if prefetch_fields:
queryset = queryset.prefetch_related(*prefetch_fields)
yield from queryset.iterator(chunk_size=self.query_batch_size)
def iter_loaded_objects(self):
grouped = self._group_by_model(self.objects.values())
for model in sorted(grouped, key=lambda item: item._meta.label_lower):
yield from self._load_objects(model, grouped[model])
def _related_objects(self, obj, custom_fields):
related_objects = []
for relation_field in obj._meta.concrete_fields:
if isinstance(relation_field, (models.ForeignKey, models.OneToOneField)):
related = getattr(obj, relation_field.name, None)
if related is not None:
related_objects.append(related)
for generic_field in obj._meta.private_fields:
if isinstance(generic_field, GenericForeignKey):
related = getattr(obj, generic_field.name, None)
if related is not None:
related_objects.append(related)
for many_to_many_field in obj._meta.many_to_many:
try:
related_objects.extend(getattr(obj, many_to_many_field.name).all())
except (AttributeError, TypeError):
pass
related_objects.extend(custom_fields)
return related_objects
def _custom_value_dependencies(self, objects, custom_fields):
target_ids_by_model = defaultdict(set)
for custom_field in custom_fields:
if custom_field.type not in ("object", "multiobject") or not custom_field.related_object_type:
continue
target_model = custom_field.related_object_type.model_class()
for obj in objects:
raw_value = obj.custom_field_data.get(custom_field.name)
if raw_value in (None, "", []):
continue
values = raw_value if custom_field.type == "multiobject" else [raw_value]
target_ids_by_model[target_model].update(values)
for obj in objects:
if ( if (
obj._meta.label_lower == "extras.customfield" obj._meta.label_lower == "extras.customfield"
and obj.type in ("object", "multiobject") and obj.type in ("object", "multiobject")
@@ -208,9 +271,37 @@ class ObjectGraph:
and obj.default not in (None, "", []) and obj.default not in (None, "", [])
): ):
target_model = obj.related_object_type.model_class() target_model = obj.related_object_type.model_class()
target_ids = obj.default if obj.type == "multiobject" else [obj.default] values = obj.default if obj.type == "multiobject" else [obj.default]
related_objects.extend(target_model._default_manager.filter(pk__in=target_ids)) target_ids_by_model[target_model].update(values)
for related in related_objects:
if self.add_dependency(related): for target_model, target_ids in target_ids_by_model.items():
dependency_queue.append(related) for target_id_batch in _chunks(target_ids, self.query_batch_size):
yield from target_model._default_manager.filter(pk__in=target_id_batch).iterator(
chunk_size=self.query_batch_size
)
def _collect_dependencies(self):
pending = dict(self.members)
scanned = set()
while pending:
next_pending = {}
for model, objects in self._group_by_model(pending.values()).items():
custom_fields = self.custom_fields_for_model(model)
loaded_objects = list(self._load_objects(model, objects))
for related in self._custom_value_dependencies(loaded_objects, custom_fields):
if self.add_dependency(related):
next_pending[object_key(related)] = related
for obj in loaded_objects:
key = object_key(obj)
if key in scanned:
continue
scanned.add(key)
for related in self._related_objects(obj, custom_fields):
if self.add_dependency(related):
next_pending[object_key(related)] = related
pending = next_pending
def collect(self, seeds):
self._expand_members(seeds)
self._collect_dependencies()
return self return self
+1
View File
@@ -48,6 +48,7 @@ class DashboardView(UserPassesTestMixin, View):
form.cleaned_data["scope_type"], form.cleaned_data["scope_type"],
scope.pk, scope.pk,
max_objects=int(get_plugin_config("netbox_export", "max_objects", 50000)), max_objects=int(get_plugin_config("netbox_export", "max_objects", 50000)),
query_batch_size=int(get_plugin_config("netbox_export", "query_batch_size", 500)),
signing_key=get_plugin_config("netbox_export", "archive_signing_key", ""), signing_key=get_plugin_config("netbox_export", "archive_signing_key", ""),
) )
except ExportImportError as exc: except ExportImportError as exc:
+1 -1
View File
@@ -4,7 +4,7 @@ build-backend = "setuptools.build_meta"
[project] [project]
name = "netbox-export" name = "netbox-export"
version = "0.2.0" version = "0.3.0"
description = "Portable ZIP export and import for scoped NetBox data" description = "Portable ZIP export and import for scoped NetBox data"
readme = "README.md" readme = "README.md"
requires-python = ">=3.12" requires-python = ">=3.12"
+1
View File
@@ -5,6 +5,7 @@ from django.conf import settings
if not settings.configured: if not settings.configured:
settings.configure( settings.configure(
DATABASES={"default": {"ENGINE": "django.db.backends.sqlite3", "NAME": ":memory:"}},
INSTALLED_APPS=["django.contrib.contenttypes"], INSTALLED_APPS=["django.contrib.contenttypes"],
SECRET_KEY="tests", SECRET_KEY="tests",
) )
+67
View File
@@ -0,0 +1,67 @@
import pytest
from django.db import connection, models
from django.test.utils import CaptureQueriesContext
from netbox_export.services import graph as graph_module
pytestmark = pytest.mark.django_db(transaction=True)
class GraphParent(models.Model):
name = models.CharField(max_length=50)
class Meta:
app_label = "graph_tests"
class GraphReference(models.Model):
code = models.CharField(max_length=50, unique=True)
class Meta:
app_label = "graph_tests"
class GraphChild(models.Model):
parent = models.ForeignKey(GraphParent, on_delete=models.CASCADE)
reference = models.ForeignKey(GraphReference, on_delete=models.PROTECT)
class Meta:
app_label = "graph_tests"
class GraphDetail(models.Model):
child = models.ForeignKey(GraphChild, on_delete=models.CASCADE)
class Meta:
app_label = "graph_tests"
def test_batched_graph_collects_members_and_dependencies(monkeypatch):
graph_models = (GraphParent, GraphReference, GraphChild, GraphDetail)
with connection.schema_editor() as schema_editor:
for model in graph_models:
schema_editor.create_model(model)
try:
reference = GraphReference.objects.create(code="shared")
parent = GraphParent.objects.create(name="scope")
children = [GraphChild(parent=parent, reference=reference) for _ in range(20)]
GraphChild.objects.bulk_create(children)
details = [GraphDetail(child=child) for child in children]
GraphDetail.objects.bulk_create(details)
monkeypatch.setattr(graph_module, "exportable_models", lambda: graph_models)
with CaptureQueriesContext(connection) as queries:
graph = graph_module.ObjectGraph(max_objects=100, query_batch_size=100).collect([parent])
loaded_keys = {graph_module.object_key(obj) for obj in graph.iter_loaded_objects()}
assert len(graph.members) == 41
assert graph_module.object_key(parent) in graph.members
assert all(graph_module.object_key(child) in graph.members for child in children)
assert all(graph_module.object_key(detail) in graph.members for detail in details)
assert set(graph.dependencies) == {graph_module.object_key(reference)}
assert loaded_keys == set(graph.objects)
assert len(queries) < 30
finally:
with connection.schema_editor() as schema_editor:
for model in reversed(graph_models):
schema_editor.delete_model(model)