Coverage for src/backend/InvenTree/plugin/base/supplier/mixins.py: 35%
69 statements
« prev ^ index » next coverage.py v7.15.2, created at 2026-10-07 17:47 +0000
« prev ^ index » next coverage.py v7.15.2, created at 2026-10-07 17:47 +0000
1"""Plugin mixin class for Supplier Integration."""
3import io
4from typing import Any, Generic, Optional, TypeVar
6import django.contrib.auth.models
7from django.core.exceptions import ValidationError
8from django.core.files.base import ContentFile
10import company.models
11import part.models as part_models
12from InvenTree.helpers_model import download_image_from_url
13from plugin import PluginMixinEnum
14from plugin.base.supplier import helpers as supplier
15from plugin.mixins import SettingsMixin
17PartData = TypeVar('PartData')
20class SupplierMixin(SettingsMixin, Generic[PartData]):
21 """Mixin which provides integration to specific suppliers."""
23 class MixinMeta:
24 """Meta options for this mixin."""
26 MIXIN_NAME = 'Supplier'
28 def __init__(self):
29 """Register mixin."""
30 super().__init__()
31 self.add_mixin(PluginMixinEnum.SUPPLIER, True, __class__)
33 self.SETTINGS['SUPPLIER'] = {
34 'name': 'Supplier',
35 'description': 'The Supplier which this plugin integrates with.',
36 'model': 'company.company',
37 'model_filters': {'is_supplier': True},
38 'required': True,
39 }
41 @property
42 def supplier_company(self):
43 """Return the supplier company object."""
44 pk = self.get_setting('SUPPLIER', cache=True)
45 if not pk:
46 raise supplier.PartImportError('Supplier setting is missing.')
48 return company.models.Company.objects.get(pk=pk)
50 # --- Methods to be overridden by plugins ---
51 def get_suppliers(self) -> list[supplier.Supplier]:
52 """Return a list of available suppliers."""
53 raise NotImplementedError('This method needs to be overridden.')
55 def get_search_results(
56 self, supplier_slug: str, term: str
57 ) -> list[supplier.SearchResult]:
58 """Return a list of search results for the given search term."""
59 raise NotImplementedError('This method needs to be overridden.')
61 def get_import_data(self, supplier_slug: str, part_id: str) -> PartData:
62 """Return the import data for the given part ID."""
63 raise NotImplementedError('This method needs to be overridden.')
65 def get_pricing_data(self, data: PartData) -> dict[int, tuple[float, str]]:
66 """Return a dictionary of pricing data for the given part data."""
67 raise NotImplementedError('This method needs to be overridden.')
69 def get_parameters(self, data: PartData) -> list[supplier.ImportParameter]:
70 """Return a list of parameters for the given part data."""
71 raise NotImplementedError('This method needs to be overridden.')
73 def import_part(
74 self,
75 data: PartData,
76 *,
77 category: Optional[part_models.PartCategory],
78 creation_user: Optional[django.contrib.auth.models.User],
79 ) -> part_models.Part:
80 """Import a part using the provided data.
82 This may include:
83 - Creating a new part
84 - Add an image to the part
85 - if this part has several variants, (create) a template part and assign it to the part
86 - create related parts
87 - add attachments to the part
88 """
89 raise NotImplementedError('This method needs to be overridden.')
91 def import_manufacturer_part(
92 self, data: PartData, *, part: part_models.Part
93 ) -> company.models.ManufacturerPart:
94 """Import a manufacturer part using the provided data.
96 This may include:
97 - Creating a new manufacturer
98 - Creating a new manufacturer part
99 - Assigning the part to the manufacturer part
100 - Setting the default supplier for the part
101 - Adding parameters to the manufacturer part
102 - Adding attachments to the manufacturer part
103 """
104 raise NotImplementedError('This method needs to be overridden.')
106 def import_supplier_part(
107 self,
108 data: PartData,
109 *,
110 part: part_models.Part,
111 manufacturer_part: company.models.ManufacturerPart,
112 ) -> company.models.SupplierPart:
113 """Import a SupplierPart using the provided data.
115 This may include:
116 - Creating a new supplier part
117 - Creating supplier price breaks
118 """
119 raise NotImplementedError('This method needs to be overridden.')
121 # --- Helper methods for importing parts ---
122 def download_image(self, img_url: str):
123 """Download an image from the given URL and return it as a ContentFile."""
124 img_r = download_image_from_url(img_url)
125 fmt = img_r.format or 'PNG'
126 buffer = io.BytesIO()
127 img_r.save(buffer, format=fmt)
129 return ContentFile(buffer.getvalue()), fmt
131 def get_template_part(
132 self, other_variants: list[part_models.Part], template_kwargs: dict[str, Any]
133 ) -> part_models.Part:
134 """Helper function to handle variant parts.
136 This helper function identifies all roots for the provided 'other_variants' list
137 - for no root => root part will be created using the 'template_kwargs'
138 - for one root
139 - root is a template => return it
140 - root is no template, create a new template like if there is no root
141 and assign it to only root that was found and return it
142 - for multiple roots => error raised
143 """
144 root_set = {v.get_root() for v in other_variants}
146 # check how much roots for the variant parts exist to identify the parent_part
147 parent_part = None # part that should be used as parent_part
148 root_part = None # part that was discovered as root part in root_set
149 if len(root_set) == 1:
150 root_part = next(iter(root_set))
151 if root_part.is_template:
152 parent_part = root_part
154 if len(root_set) == 0 or (root_part and not root_part.is_template):
155 parent_part = part_models.Part.objects.create(**template_kwargs)
157 if not parent_part:
158 raise supplier.PartImportError(
159 f'A few variant parts from the supplier are already imported, but have different InvenTree variant root parts, try to merge them to the same root variant template part (parts: {", ".join(str(p.pk) for p in other_variants)}).'
160 )
162 # assign parent_part to root_part if root_part has no variant of already
163 if root_part and not root_part.is_template and not root_part.variant_of:
164 root_part.variant_of = parent_part
165 root_part.save()
167 return parent_part
169 def create_related_parts(
170 self, part: part_models.Part, related_parts: list[part_models.Part]
171 ):
172 """Create relationships between the given part and related parts."""
173 for p in related_parts:
174 try:
175 part_models.PartRelated.objects.create(part_1=part, part_2=p)
176 except ValidationError:
177 pass # pass, duplicate relationship detected