Coverage for utilities/forms/utils.py: 11%
154 statements
« prev ^ index » next coverage.py v7.15.2, created at 2026-10-10 18:35 +0000
« prev ^ index » next coverage.py v7.15.2, created at 2026-10-10 18:35 +0000
1import re
3from django import forms
4from django.forms.models import fields_for_model
5from django.utils.translation import gettext as _
7from utilities.choices import unpack_grouped_choices
8from utilities.querysets import RestrictedQuerySet
10from .constants import *
12__all__ = (
13 'add_blank_choice',
14 'expand_alphanumeric_pattern',
15 'expand_ipnetwork_pattern',
16 'form_from_model',
17 'get_capacity_unit_label',
18 'get_field_value',
19 'get_selected_values',
20 'parse_alphanumeric_range',
21 'parse_csv',
22 'parse_numeric_range',
23 'restrict_form_fields',
24 'validate_csv',
25)
28def parse_numeric_range(string, base=10, min_value=None, max_value=None):
29 """
30 Expand a numeric range (continuous or not) into a decimal or
31 hexadecimal list, as specified by the base parameter
32 '0-3,5' => [0, 1, 2, 3, 5]
33 '2,8-b,d,f' => [2, 8, 9, a, b, d, f]
35 Pass BOTH ``min_value`` and ``max_value`` to validate each range against those bounds *before* it is
36 expanded: a reversed or out-of-bounds range then raises rather than materializing a huge list or
37 silently expanding to nothing (which would be swallowed when combined with valid ranges, e.g.
38 "80,9000-53"). Bounds are all-or-nothing — supplying only one raises ``ValueError`` — so a caller
39 can't opt into a lower bound while leaving the expansion size uncapped. With no bounds (e.g.
40 IP/pattern expansion) a reversed range yields an empty list, as before.
41 """
42 bounded = min_value is not None or max_value is not None
43 if bounded and (min_value is None or max_value is None):
44 raise ValueError("parse_numeric_range() requires both min_value and max_value, or neither.")
46 values = list()
47 for dash_range in string.split(','):
48 try:
49 begin, end = dash_range.split('-')
50 except ValueError:
51 begin, end = dash_range, dash_range
52 try:
53 begin, end = int(begin.strip(), base=base), int(end.strip(), base=base) + 1
54 except ValueError:
55 raise forms.ValidationError(_('Range "{value}" is invalid.').format(value=dash_range))
56 if bounded:
57 # Reject reversed ranges and endpoints outside the permitted range before expanding.
58 if begin > end - 1:
59 raise forms.ValidationError(_('Range "{value}" is invalid.').format(value=dash_range))
60 if begin < min_value or end - 1 > max_value:
61 raise forms.ValidationError(
62 _('Range "{value}" is not within the permitted range ({min}-{max}).').format(
63 value=dash_range, min=min_value, max=max_value
64 )
65 )
66 values.extend(range(begin, end))
67 return sorted(set(values))
70def parse_alphanumeric_range(string):
71 """
72 Expand an alphanumeric range (continuous or not) into a list.
73 'a-d,f' => [a, b, c, d, f]
74 '0-3,a-d' => [0, 1, 2, 3, a, b, c, d]
75 """
76 values = []
77 for value in string.split(','):
78 if '-' not in value:
79 # Item is not a range
80 values.append(value)
81 continue
83 # Find the range's beginning & end values
84 try:
85 begin, end = value.split('-')
86 vals = begin + end
87 # Break out of loop if there's an invalid pattern to return an error
88 if (not (vals.isdigit() or vals.isalpha())) or (vals.isalpha() and not (vals.isupper() or vals.islower())):
89 return []
90 except ValueError:
91 raise forms.ValidationError(_('Range "{value}" is invalid.').format(value=value))
93 # Numeric range
94 if begin.isdigit() and end.isdigit():
95 if int(begin) >= int(end):
96 raise forms.ValidationError(
97 _('Invalid range: Ending value ({end}) must be greater than beginning value ({begin}).').format(
98 begin=begin, end=end
99 )
100 )
101 for n in list(range(int(begin), int(end) + 1)):
102 values.append(n)
104 # Alphanumeric range
105 else:
106 # Not a valid range (more than a single character)
107 if not len(begin) == len(end) == 1:
108 raise forms.ValidationError(_('Range "{value}" is invalid.').format(value=value))
109 if ord(begin) >= ord(end):
110 raise forms.ValidationError(_('Range "{value}" is invalid.').format(value=value))
111 for n in list(range(ord(begin), ord(end) + 1)):
112 values.append(chr(n))
114 return values
117def expand_alphanumeric_pattern(string):
118 """
119 Expand an alphabetic pattern into a list of strings.
120 """
121 lead, pattern, remnant = re.split(ALPHANUMERIC_EXPANSION_PATTERN, string, maxsplit=1)
122 parsed_range = parse_alphanumeric_range(pattern)
123 for i in parsed_range:
124 if re.search(ALPHANUMERIC_EXPANSION_PATTERN, remnant):
125 for string in expand_alphanumeric_pattern(remnant):
126 yield "{}{}{}".format(lead, i, string)
127 else:
128 yield "{}{}{}".format(lead, i, remnant)
131def expand_ipnetwork_pattern(string, family):
132 """
133 Expand an IP network pattern into a list of strings. Examples:
134 '192.0.2.[1,2,100-250]/24' => ['192.0.2.1/24', '192.0.2.2/24', '192.0.2.100/24' ... '192.0.2.250/24']
135 '2001:db8:0:[0,fd-ff]::/64' => ['2001:db8:0:0::/64', '2001:db8:0:fd::/64', ... '2001:db8:0:ff::/64']
136 """
137 if family not in [4, 6]:
138 raise Exception("Invalid IP address family: {}".format(family))
139 if family == 4:
140 regex = IP4_EXPANSION_PATTERN
141 base = 10
142 else:
143 regex = IP6_EXPANSION_PATTERN
144 base = 16
145 lead, pattern, remnant = re.split(regex, string, maxsplit=1)
146 parsed_range = parse_numeric_range(pattern, base)
147 for i in parsed_range:
148 if re.search(regex, remnant):
149 for string in expand_ipnetwork_pattern(remnant, family):
150 yield ''.join([lead, format(i, 'x' if family == 6 else 'd'), string])
151 else:
152 yield ''.join([lead, format(i, 'x' if family == 6 else 'd'), remnant])
155def get_capacity_unit_label(divisor=1000):
156 """
157 Return the appropriate base unit label: 'MiB' for binary (1024), 'MB' for decimal (1000).
158 """
159 return 'MiB' if divisor == 1024 else 'MB'
162def get_field_value(form, field_name):
163 """
164 Return the current bound or initial value associated with a form field, prior to calling
165 clean() for the form.
166 """
167 field = form.fields[field_name]
169 if form.is_bound and field_name in form.data:
170 if (value := form.data[field_name]) is None:
171 return None
172 if hasattr(field, 'valid_value') and field.valid_value(value):
173 return value
175 return form.get_initial_for_field(field, field_name)
178def get_selected_values(form, field_name):
179 """
180 Return the list of selected human-friendly values for a form field
181 """
182 if not hasattr(form, 'cleaned_data'):
183 form.is_valid()
184 filter_data = form.cleaned_data.get(field_name)
185 field = form.fields[field_name]
187 # Non-selection field
188 if not hasattr(field, 'choices'):
189 return [str(filter_data)]
191 # Model choice field
192 if type(field.choices) is forms.models.ModelChoiceIterator:
193 # If this is a single-choice field, wrap its value in a list
194 if not hasattr(filter_data, '__iter__'):
195 values = [filter_data]
196 else:
197 values = filter_data
199 else:
200 # Static selection field
201 choices = unpack_grouped_choices(field.choices)
202 if type(filter_data) not in (list, tuple):
203 filter_data = [filter_data] # Ensure filter data is iterable
204 values = [
205 label for value, label in choices if str(value) in filter_data or None in filter_data
206 ]
208 # If the field has a `null_option` attribute set and it is selected,
209 # add it to the field's grouped choices.
210 if getattr(field, 'null_option', None) and None in filter_data:
211 values.remove(None)
212 values.insert(0, field.null_option)
214 return values
217def add_blank_choice(choices):
218 """
219 Add a blank choice to the beginning of a choices list. Any Choice objects are preserved (rather than reduced to
220 plain tuples) so that description-aware fields can still reference their descriptions.
221 """
222 return ((None, '---------'), *choices)
225def form_from_model(model, fields):
226 """
227 Return a Form class with the specified fields derived from a model. This is useful when we need a form to be used
228 for creating objects, but want to avoid the model's validation (e.g. for bulk create/edit functions). All fields
229 are marked as not required.
230 """
231 form_fields = fields_for_model(model, fields=fields)
232 for field in form_fields.values():
233 field.required = False
234 field.widget.is_required = False
236 return type('FormFromModel', (forms.Form,), form_fields)
239def restrict_form_fields(form, user, action='view'):
240 """
241 Restrict all form fields which reference a RestrictedQuerySet. This ensures that users see only permitted objects
242 as available choices.
243 """
244 for field in form.fields.values():
245 if hasattr(field, 'queryset') and issubclass(field.queryset.__class__, RestrictedQuerySet):
246 field.queryset = field.queryset.restrict(user, action)
249def parse_csv(reader):
250 """
251 Parse a csv_reader object into a headers dictionary and a list of records dictionaries. Raise an error
252 if the records are formatted incorrectly. Return headers and records as a tuple.
253 """
254 records = []
255 headers = {}
257 # Consume the first line of CSV data as column headers. Create a dictionary mapping each header to an optional
258 # "to" field specifying how the related object is being referenced. For example, importing a Device might use a
259 # `site.slug` header, to indicate the related site is being referenced by its slug.
261 for header in next(reader):
262 header = header.strip()
263 if '.' in header:
264 field, to_field = header.split('.', 1)
265 if field in headers:
266 raise forms.ValidationError(_('Duplicate or conflicting column header for "{field}"').format(
267 field=field
268 ))
269 headers[field] = to_field
270 else:
271 if header in headers:
272 raise forms.ValidationError(_('Duplicate or conflicting column header for "{header}"').format(
273 header=header
274 ))
275 headers[header] = None
277 # Parse CSV rows into a list of dictionaries mapped from the column headers.
278 for i, row in enumerate(reader, start=1):
279 if len(row) != len(headers):
280 raise forms.ValidationError(
281 _("Row {row}: Expected {count_expected} columns but found {count_found}").format(
282 row=i, count_expected=len(headers), count_found=len(row)
283 )
284 )
285 row = [col.strip() for col in row]
286 record = dict(zip(headers.keys(), row))
287 records.append(record)
289 return headers, records
292def validate_csv(headers, fields, required_fields):
293 """
294 Validate that parsed csv data conforms to the object's available fields. Raise validation errors
295 if parsed csv data contains invalid headers or does not contain required headers.
296 """
297 # Validate provided column headers
298 is_update = False
299 for field, to_field in headers.items():
300 if field == "id":
301 is_update = True
302 continue
303 if field not in fields:
304 raise forms.ValidationError(_('Unexpected column header "{field}" found.').format(field=field))
305 if to_field and not hasattr(fields[field], 'to_field_name'):
306 raise forms.ValidationError(_('Column "{field}" is not a related object; cannot use dots').format(
307 field=field
308 ))
309 if to_field and not hasattr(fields[field].queryset.model, to_field):
310 raise forms.ValidationError(_('Invalid related object attribute for column "{field}": {to_field}').format(
311 field=field, to_field=to_field
312 ))
314 # Validate required fields (if not an update)
315 if not is_update:
316 for f in required_fields:
317 if f not in headers:
318 raise forms.ValidationError(_('Required column header "{header}" not found.').format(header=f))