Files
NetBox-Export/netbox_export/services/graph.py
T

308 lines
13 KiB
Python

from __future__ import annotations
from collections import defaultdict
from django.apps import apps
from django.contrib.contenttypes.fields import GenericForeignKey
from django.contrib.contenttypes.models import ContentType
from django.db import models
from .codec import object_key
from .exceptions import GraphLimitError
EXCLUDED_APP_LABELS = {
"account",
"admin",
"auth",
"contenttypes",
"sessions",
"users",
"netbox_export",
}
EXCLUDED_MODELS = {
"core.job",
"core.objectchange",
"extras.eventrule",
"extras.journalentry",
"extras.notification",
"extras.notificationgroup",
"extras.savedfilter",
"extras.subscription",
}
SCOPE_LINK_FIELDS = {"tenant", "site", "location", "region"}
PEER_CONTAINER_MODELS = {"dcim.cable", "circuits.circuit", "circuits.virtualcircuit"}
def is_exportable_model(model) -> bool:
opts = model._meta
return bool(
opts.managed
and not opts.abstract
and not opts.proxy
and not opts.auto_created
and opts.app_label not in EXCLUDED_APP_LABELS
and opts.label_lower not in EXCLUDED_MODELS
)
def exportable_models():
return tuple(model for model in apps.get_models() if is_exportable_model(model))
def _descendants(obj):
if hasattr(obj, "get_descendants"):
return list(obj.get_descendants(include_self=True))
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):
from dcim.models import Location, Region, Site
from tenancy.models import Tenant, TenantGroup
models_by_scope = {
"tenant_group": TenantGroup,
"tenant": Tenant,
"region": Region,
"site": Site,
"location": Location,
}
model = models_by_scope[scope_type]
root = model.objects.get(pk=scope_id)
seeds = _descendants(root)
if scope_type == "tenant_group":
group_ids = [obj.pk for obj in seeds]
seeds.extend(Tenant.objects.filter(group_id__in=group_ids))
elif scope_type == "region":
region_ids = [obj.pk for obj in seeds]
seeds.extend(Site.objects.filter(region_id__in=region_ids))
return root, seeds
class ObjectGraph:
"""Collect scoped objects and dependencies using batched relation queries."""
def __init__(self, max_objects: int, query_batch_size: int = 500):
self.max_objects = max_objects
self.query_batch_size = max(1, query_batch_size)
self.members: 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
def objects(self) -> dict[str, models.Model]:
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):
object_count = len(self.members) + len(self.dependencies)
if object_count > self.max_objects:
raise GraphLimitError(
f"Der Export würde mehr als {self.max_objects} Objekte enthalten. "
"Bitte den Bereich verkleinern oder max_objects erhöhen."
)
def add_member(self, obj) -> bool:
key = object_key(obj)
if key in self.members:
return False
self.dependencies.pop(key, None)
self.members[key] = obj
self._check_limit()
return True
def add_dependency(self, obj) -> bool:
key = object_key(obj)
if key in self.members or key in self.dependencies or not is_exportable_model(type(obj)):
return False
self.dependencies[key] = obj
self._check_limit()
return True
def _group_by_model(self, objects):
grouped = defaultdict(list)
for obj in objects:
grouped[type(obj)].append(obj)
return grouped
def _expand_members(self, seeds):
pending = {}
for obj in seeds:
if self.add_member(obj):
pending[object_key(obj)] = obj
while pending:
next_pending = {}
for parent_model, parents in self._group_by_model(pending.values()).items():
parent_ct = ContentType.objects.get_for_model(parent_model)
direct_candidates = self._reverse_fields.get(parent_model, {})
peer_fields = self._peer_fields.get(parent_model, [])
for parent_batch in _chunks(parents, self.query_batch_size):
parent_ids = [obj.pk for obj in parent_batch]
if peer_fields:
peer_names = [field.name for field in peer_fields]
queryset = parent_model._default_manager.filter(pk__in=parent_ids).select_related(*peer_names)
for parent in queryset.iterator(chunk_size=self.query_batch_size):
for relation_field in peer_fields:
related = getattr(parent, relation_field.name, None)
if related is not None and self.add_member(related):
next_pending[object_key(related)] = related
candidate_queries = defaultdict(models.Q)
for candidate_model, relation_fields in direct_candidates.items():
for relation_field in relation_fields:
candidate_queries[candidate_model] |= models.Q(**{f"{relation_field.attname}__in": parent_ids})
for candidate_model, generic_fields in self._generic_fields.items():
for generic_field in generic_fields:
candidate_queries[candidate_model] |= models.Q(
**{
generic_field.ct_field: parent_ct,
f"{generic_field.fk_field}__in": parent_ids,
}
)
for candidate_model, query in candidate_queries.items():
queryset = candidate_model._default_manager.filter(query).distinct()
for child in queryset.iterator(chunk_size=self.query_batch_size):
if self.add_member(child):
next_pending[object_key(child)] = child
pending = next_pending
def custom_fields_for_model(self, model):
if model not in self._custom_field_cache:
if hasattr(model, "custom_field_data"):
from extras.models import CustomField
self._custom_field_cache[model] = list(CustomField.objects.get_for_model(model))
else:
self._custom_field_cache[model] = []
return self._custom_field_cache[model]
def _load_objects(self, model, objects):
foreign_keys = [
field.name
for field in model._meta.concrete_fields
if isinstance(field, (models.ForeignKey, models.OneToOneField))
]
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 (
obj._meta.label_lower == "extras.customfield"
and obj.type in ("object", "multiobject")
and obj.related_object_type
and obj.default not in (None, "", [])
):
target_model = obj.related_object_type.model_class()
values = obj.default if obj.type == "multiobject" else [obj.default]
target_ids_by_model[target_model].update(values)
for target_model, target_ids in target_ids_by_model.items():
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