Coverage for src/backend/InvenTree/InvenTree/helpers.py: 38%
518 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"""Provides helper functions used throughout the InvenTree project."""
3import datetime
4import hashlib
5import inspect
6import io
7import json
8import os.path
9import re
10from decimal import Decimal, InvalidOperation
11from pathlib import Path
12from typing import Optional, TypeVar
13from wsgiref.util import FileWrapper
14from zoneinfo import ZoneInfo, ZoneInfoNotFoundError
16from django.conf import settings
17from django.contrib.staticfiles.storage import StaticFilesStorage
18from django.core.exceptions import FieldError, ValidationError
19from django.core.files.storage import default_storage
20from django.db.models.fields.files import FieldFile, ImageFieldFile
21from django.http import StreamingHttpResponse
22from django.utils import timezone
23from django.utils.translation import gettext_lazy as _
25import nh3
26import structlog
27from djmoney.money import Money
28from PIL import Image
29from stdimage.models import StdImageField, StdImageFieldFile
31from common.currency import currency_code_default
32from InvenTree.sanitizer import (
33 DEAFAULT_ATTRS,
34 DEFAULT_CSS,
35 DEFAULT_PROTOCOLS,
36 DEFAULT_TAGS,
37)
39logger = structlog.get_logger('inventree')
41INT_CLIP_MAX = 0x7FFFFFFF
44def extract_int(
45 reference, clip=INT_CLIP_MAX, try_hex=False, allow_negative=False
46) -> int:
47 """Extract an integer out of provided string.
49 Arguments:
50 reference: Input string to extract integer from
51 clip: Maximum value to return (default = 0x7FFFFFFF)
52 try_hex: Attempt to parse as hex if integer conversion fails (default = False)
53 allow_negative: Allow negative values (default = False)
54 """
55 # Default value if we cannot convert to an integer
56 ref_int = 0
58 def do_clip(value: int, clip: int, allow_negative: bool) -> int:
59 """Perform clipping on the provided value.
61 Arguments:
62 value: Value to clip
63 clip: Maximum value to clip to
64 allow_negative: Allow negative values (default = False)
65 """
66 if clip is None: 66 ↛ 67line 66 didn't jump to line 67 because the condition on line 66 was never true
67 return value
69 clip = min(clip, INT_CLIP_MAX)
71 if value > clip: 71 ↛ 72line 71 didn't jump to line 72 because the condition on line 71 was never true
72 return clip
73 elif value < -clip: 73 ↛ 74line 73 didn't jump to line 74 because the condition on line 73 was never true
74 return -clip
76 if not allow_negative: 76 ↛ 79line 76 didn't jump to line 79 because the condition on line 76 was always true
77 value = abs(value)
79 return value
81 reference = str(reference).strip()
83 # Ignore empty string
84 if len(reference) == 0: 84 ↛ 85line 84 didn't jump to line 85 because the condition on line 84 was never true
85 return 0
87 # Try naive integer conversion first
88 try:
89 ref_int = int(reference)
90 return do_clip(ref_int, clip, allow_negative)
91 except ValueError:
92 pass
94 # Hex?
95 if try_hex or reference.startswith('0x'):
96 try:
97 ref_int = int(reference, base=16)
98 return do_clip(ref_int, clip, allow_negative)
99 except ValueError:
100 pass
102 # Look at the start of the string - can it be "integerized"?
103 result = re.match(r'^(\d+)', reference)
105 if result and len(result.groups()) == 1:
106 ref = result.groups()[0]
107 try:
108 ref_int = int(ref)
109 except Exception:
110 ref_int = 0
111 else:
112 # Look at the "end" of the string
113 result = re.search(r'(\d+)$', reference)
115 if result and len(result.groups()) == 1:
116 ref = result.groups()[0]
117 try:
118 ref_int = int(ref)
119 except Exception:
120 ref_int = 0
122 # Ensure that the returned values are within the range that can be stored in an IntegerField
123 # Note: This will result in large values being "clipped"
124 ref_int = do_clip(ref_int, clip, allow_negative)
126 if not allow_negative and ref_int < 0:
127 ref_int = abs(ref_int)
129 return ref_int
132def generateTestKey(test_name: str | None) -> str:
133 """Generate a test 'key' for a given test name. This must not have illegal chars as it will be used for dict lookup in a template.
135 Tests must be named such that they will have unique keys.
136 """
137 if test_name is None: 137 ↛ 138line 137 didn't jump to line 138 because the condition on line 137 was never true
138 test_name = ''
140 key = test_name.strip().lower()
141 key = key.replace(' ', '')
143 def valid_char(char: str):
144 """Determine if a particular character is valid for use in a test key."""
145 if not char.isprintable(): 145 ↛ 148line 145 didn't jump to line 148 because the condition on line 145 was always true
146 return False
148 if char.isidentifier():
149 return True
151 return bool(char.isalnum())
153 # Remove any characters that cannot be used to represent a variable
154 key = ''.join([c for c in key if valid_char(c)])
156 # If the key starts with a non-identifier character, prefix with an underscore
157 if len(key) > 0 and not key[0].isidentifier(): 157 ↛ 158line 157 didn't jump to line 158 because the condition on line 157 was never true
158 key = '_' + key
160 return key
163def constructPathString(path: list[str], max_chars: int = 250) -> str:
164 """Construct a 'path string' for the given path.
166 Arguments:
167 path: A list of strings e.g. ['path', 'to', 'location']
168 max_chars: Maximum number of characters
169 """
170 pathstring = '/'.join(path)
172 # Replace middle elements to limit the pathstring
173 if len(pathstring) > max_chars: 173 ↛ 174line 173 didn't jump to line 174 because the condition on line 173 was never true
174 n = int(max_chars / 2 - 2)
175 pathstring = pathstring[:n] + '...' + pathstring[-n:]
177 return pathstring
180def getMediaUrl(
181 file: FieldFile | ImageFieldFile | StdImageFieldFile, name: str | None = None
182):
183 """Return the qualified access path for the given file, under the media directory."""
184 if not isinstance(file, (FieldFile, ImageFieldFile, StdImageFieldFile)):
185 raise TypeError(
186 'file must be one of FileField, ImageFileField, StdImageFieldFile'
187 )
188 if name is not None:
189 file = regenerate_imagefile(file, name)
191 return default_storage.url(file.name)
194def regenerate_imagefile(_file, _name: str):
195 """Regenerate a file object for a given variation name.
197 Arguments:
198 _file: Original file object
199 _name: Name of the variation (e.g. 'thumbnail', 'preview')
200 """
201 name = _file.field.attr_class.get_variation_name(_file.name, _name)
202 return ImageFieldFile(_file.instance, _file, name) # ty:ignore[too-many-positional-arguments]
205def image2name(img_obj: StdImageField, do_preview: bool, do_thumbnail: bool):
206 """Convert an image object to a filename string.
208 Arguments:
209 img_obj: Image object
210 do_preview: Return preview image name
211 do_thumbnail: Return thumbnail image name
212 """
214 def extract(ref: str):
215 return None if not hasattr(img_obj, ref) else getattr(img_obj, ref).name
217 if not img_obj:
218 return None
219 elif do_preview:
220 return extract('preview')
221 elif do_thumbnail:
222 return extract('thumbnail')
223 else:
224 return img_obj.name
227def getStaticUrl(filename):
228 """Return the qualified access path for the given file, under the static media directory."""
229 return StaticFilesStorage().url(filename)
232def TestIfImage(img) -> bool:
233 """Test if an image file is indeed an image.
235 Arguments:
236 img: A file-like object
238 Returns:
239 True if the file is a valid image, False otherwise
240 """
241 try:
242 Image.open(img).verify()
243 return True
244 except Exception:
245 return False
248def getBlankImage():
249 """Return the qualified path for the 'blank image' placeholder."""
250 return getStaticUrl('img/blank_image.png')
253def getBlankThumbnail():
254 """Return the qualified path for the 'blank image' thumbnail placeholder."""
255 return getStaticUrl('img/blank_image.thumbnail.png')
258def checkStaticFile(*args) -> bool:
259 """Check if a file exists in the static storage."""
260 static_storage = StaticFilesStorage()
261 fn = Path(*args)
262 return static_storage.exists(str(fn))
265def getLogoImage(as_file=False, custom=True):
266 """Return the InvenTree logo image, or a custom logo if available."""
267 if custom and settings.CUSTOM_LOGO: 267 ↛ 283line 267 didn't jump to line 283 because the condition on line 267 was always true
268 static_storage = StaticFilesStorage()
270 if static_storage.exists(settings.CUSTOM_LOGO): 270 ↛ 271line 270 didn't jump to line 271 because the condition on line 270 was never true
271 storage = static_storage
272 elif default_storage.exists(settings.CUSTOM_LOGO): 272 ↛ 273line 272 didn't jump to line 273 because the condition on line 272 was never true
273 storage = default_storage
274 else:
275 storage = None
277 if storage is not None: 277 ↛ 278line 277 didn't jump to line 278 because the condition on line 277 was never true
278 if as_file:
279 return f'file://{storage.path(settings.CUSTOM_LOGO)}'
280 return storage.url(settings.CUSTOM_LOGO)
282 # If we have got to this point, return the default logo
283 if as_file: 283 ↛ 284line 283 didn't jump to line 284 because the condition on line 283 was never true
284 path = settings.STATIC_ROOT.joinpath('img/inventree.png')
285 return f'file://{path}'
286 return getStaticUrl('img/inventree.png')
289def getSplashScreen(custom=True):
290 """Return the InvenTree splash screen, or a custom splash if available."""
291 static_storage = StaticFilesStorage()
293 if custom and settings.CUSTOM_SPLASH: 293 ↛ 298line 293 didn't jump to line 298 because the condition on line 293 was always true
294 if static_storage.exists(settings.CUSTOM_SPLASH): 294 ↛ 295line 294 didn't jump to line 295 because the condition on line 294 was never true
295 return static_storage.url(settings.CUSTOM_SPLASH)
297 # No custom splash screen
298 return static_storage.url('img/inventree_splash.jpg')
301def getCustomOption(reference: str):
302 """Return the value of a custom option from settings.CUSTOMIZE.
304 Args:
305 reference: Reference key for the custom option
306 """
307 return settings.CUSTOMIZE.get(reference, None)
310def TestIfImageURL(url):
311 """Test if an image URL (or filename) looks like a valid image format.
313 Simply tests the extension against a set of allowed values
314 """
315 return os.path.splitext(os.path.basename(url))[-1].lower() in [
316 '.jpg',
317 '.jpeg',
318 '.j2k',
319 '.png',
320 '.bmp',
321 '.tif',
322 '.tiff',
323 '.webp',
324 '.gif',
325 ]
328def str2bool(text, test=True) -> bool:
329 """Test if a string 'looks' like a boolean value.
331 Args:
332 text: Input text
333 test (default = True): Set which boolean value to look for
335 Returns:
336 True if the text looks like the selected boolean value
337 """
338 if test:
339 return str(text).lower() in ['1', 'y', 'yes', 't', 'true', 'ok', 'on']
340 return str(text).lower() in ['0', 'n', 'no', 'none', 'f', 'false', 'off']
343def is_bool(text: str) -> bool:
344 """Determine if a string value 'looks' like a boolean."""
345 return str2bool(text, True) or str2bool(text, False)
348def isNull(text: str) -> bool:
349 """Test if a string 'looks' like a null value. This is useful for querying the API against a null key.
351 Args:
352 text: Input text
354 Returns:
355 True if the text looks like a null value
356 """
357 return str(text).strip().lower() in [
358 'top',
359 'null',
360 'none',
361 'empty',
362 'false',
363 '-1',
364 '',
365 ]
368def normalize(d, rounding: Optional[int] = None) -> Decimal:
369 """Normalize a decimal number, and remove exponential formatting."""
370 if type(d) is not Decimal:
371 d = Decimal(d)
373 if rounding is not None:
374 d = round(d, rounding)
376 d = d.normalize()
378 # Ref: https://docs.python.org/3/library/decimal.html
379 return d.quantize(Decimal(1)) if d == d.to_integral() else d.normalize()
382def increment(value):
383 """Attempt to increment an integer (or a string that looks like an integer).
385 e.g.
387 001 -> 002
388 2 -> 3
389 AB01 -> AB02
390 QQQ -> QQQ
392 """
393 # Ignore empty strings
394 if value in ['', None]:
395 # Provide a default value if provided with a null input
396 return '1'
398 value = str(value).strip()
400 pattern = r'(.*?)(\d+)?$'
402 result = re.search(pattern, value)
404 # No match!
405 if result is None: 405 ↛ 406line 405 didn't jump to line 406 because the condition on line 405 was never true
406 return value
408 groups = result.groups()
410 # If we cannot match the regex, then simply return the provided value
411 if len(groups) != 2: 411 ↛ 412line 411 didn't jump to line 412 because the condition on line 411 was never true
412 return value
414 prefix, number = groups
416 # No number extracted? Simply return the prefix (without incrementing!)
417 if not number: 417 ↛ 418line 417 didn't jump to line 418 because the condition on line 417 was never true
418 return prefix
420 # Record the width of the number
421 width = len(number)
423 try:
424 number = int(number) + 1
425 number = str(number)
426 except ValueError:
427 pass
429 return prefix + str(number).zfill(width)
432def decimal2string(d):
433 """Format a Decimal number as a string, stripping out any trailing zeroes or decimal points. Essentially make it look like a whole number if it is one.
435 Args:
436 d: A python Decimal object
438 Returns:
439 A string representation of the input number
440 """
441 if type(d) is Decimal:
442 d = normalize(d)
444 try:
445 # Ensure that the provided string can actually be converted to a float
446 float(d)
447 except ValueError:
448 # Not a number
449 return str(d)
451 s = str(d)
453 # Return entire number if there is no decimal place
454 if '.' not in s:
455 return s
457 return s.rstrip('0').rstrip('.')
460def decimal2money(d, currency=None):
461 """Format a Decimal number as Money.
463 Args:
464 d: A python Decimal object
465 currency: Currency of the input amount, defaults to default currency in settings
467 Returns:
468 A Money object from the input(s)
469 """
470 if not currency:
471 currency = currency_code_default()
472 return Money(d, currency)
475def WrapWithQuotes(text, quote='"'):
476 """Wrap the supplied text with quotes.
478 Args:
479 text: Input text to wrap
480 quote: Quote character to use for wrapping (default = "")
482 Returns:
483 Supplied text wrapped in quote char
484 """
485 if not text.startswith(quote):
486 text = quote + text
488 if not text.endswith(quote):
489 text = text + quote
491 return text
494def GetExportOptions() -> list:
495 """Return a set of allowable import / export file formats."""
496 return [['csv', 'CSV'], ['xlsx', 'Excel'], ['tsv', 'TSV']]
499def GetExportFormats() -> list:
500 """Return a list of allowable file formats for importing or exporting tabular data."""
501 return [opt[0] for opt in GetExportOptions()]
504def DownloadFile(
505 data, filename, content_type='application/text', inline=False
506) -> StreamingHttpResponse:
507 """Create a dynamic file for the user to download.
509 Args:
510 data: Raw file data (string or bytes)
511 filename: Filename for the file download
512 content_type: Content type for the download
513 inline: Download "inline" or as attachment? (Default = attachment)
515 Return:
516 A StreamingHttpResponse object wrapping the supplied data
517 """
518 filename = WrapWithQuotes(filename)
519 length = len(data)
521 if isinstance(data, str):
522 wrapper = FileWrapper(io.StringIO(data))
523 else:
524 wrapper = FileWrapper(io.BytesIO(data))
526 response = StreamingHttpResponse(wrapper, content_type=content_type)
527 if isinstance(data, str):
528 length = len(bytes(data, response.charset))
529 response['Content-Length'] = length
531 if inline:
532 disposition = f'inline; filename={filename}'
533 else:
534 disposition = f'attachment; filename={filename}'
536 response['Content-Disposition'] = disposition
537 return response
540def increment_serial_number(serial, part=None):
541 """Given a serial number, (attempt to) generate the *next* serial number.
543 Note: This method is exposed to custom plugins.
545 Arguments:
546 serial: The serial number which should be incremented
547 part: Optional part object to provide additional context for incrementing the serial number
549 Returns:
550 incremented value, or None if incrementing could not be performed.
551 """
552 from InvenTree.exceptions import log_error
553 from InvenTree.ready import isReadOnlyCommand
554 from plugin import PluginMixinEnum, registry
556 # Ensure we start with a string value
557 if serial is not None: 557 ↛ 558line 557 didn't jump to line 558 because the condition on line 557 was never true
558 serial = str(serial).strip()
560 if not isReadOnlyCommand(): 560 ↛ 581line 560 didn't jump to line 581 because the condition on line 560 was always true
561 # First, let any plugins attempt to increment the serial number
562 for plugin in registry.with_mixin(PluginMixinEnum.VALIDATION): 562 ↛ 563line 562 didn't jump to line 563 because the loop on line 562 never started
563 try:
564 if not hasattr(plugin, 'increment_serial_number'):
565 continue
567 signature = inspect.signature(plugin.increment_serial_number)
569 # Note: 2024-08-21 - The 'part' parameter has been added to the signature
570 if 'part' in signature.parameters:
571 result = plugin.increment_serial_number(serial, part=part)
572 else:
573 result = plugin.increment_serial_number(serial)
574 if result is not None:
575 return str(result)
576 except Exception:
577 log_error('increment_serial_number', plugin=plugin.slug)
579 # If we get to here, no plugins were able to "increment" the provided serial value
580 # Attempt to perform increment according to some basic rules
581 return increment(serial)
584def extract_serial_numbers(
585 input_string, expected_quantity: int, starting_value=None, part=None
586):
587 """Extract a list of serial numbers from a provided input string.
589 The input string can be specified using the following concepts:
591 - Individual serials are separated by comma: 1, 2, 3, 6,22
592 - Sequential ranges with provided limits are separated by hyphens: 1-5, 20 - 40
593 - The "next" available serial number can be specified with the tilde (~) character
594 - Serial numbers can be supplied as <start>+ for getting all expected numbers starting from <start>
595 - Serial numbers can be supplied as <start>+<length> for getting <length> numbers starting from <start>
597 Actual generation of sequential serials is passed to the 'validation' plugin mixin,
598 allowing custom plugins to determine how serial values are incremented.
600 Arguments:
601 input_string: Input string with specified serial numbers (string, or integer)
602 expected_quantity: The number of (unique) serial numbers we expect
603 starting_value: Provide a starting value for the sequence (or None)
604 part: Part that should be used as context
605 """
606 if starting_value is None:
607 starting_value = increment_serial_number(None, part=part)
609 try:
610 expected_quantity = int(expected_quantity)
611 except ValueError:
612 raise ValidationError([_('Invalid quantity provided')])
614 if expected_quantity > 1000:
615 raise ValidationError({
616 'quantity': [_('Cannot serialize more than 1000 items at once')]
617 })
619 input_string = str(input_string).strip() if input_string else ''
621 if len(input_string) == 0:
622 raise ValidationError([_('Empty serial number string')])
624 next_value = increment_serial_number(starting_value, part=part)
626 # Substitute ~ character with latest value
627 while '~' in input_string and next_value:
628 input_string = input_string.replace('~', str(next_value), 1)
629 next_value = increment_serial_number(next_value, part=part)
631 # Split input string by whitespace or comma (,) characters
632 groups = re.split(r'[\s,]+', input_string)
634 serials = []
635 errors = []
637 def add_error(error: str):
638 """Helper function for adding an error message."""
639 if error not in errors:
640 errors.append(error)
642 def add_serial(serial):
643 """Helper function to check for duplicated values."""
644 serial = serial.strip()
646 # Ignore blank / empty serials
647 if len(serial) == 0:
648 return
650 if serial in serials:
651 add_error(_('Duplicate serial') + f': {serial}')
652 else:
653 serials.append(serial)
655 # If the user has supplied the correct number of serials, do not split into groups
656 if len(groups) == expected_quantity:
657 for group in groups:
658 add_serial(group)
660 if len(errors) > 0:
661 raise ValidationError(errors)
662 else:
663 return serials
665 for group in groups:
666 # Calculate the "remaining" quantity of serial numbers
667 remaining = expected_quantity - len(serials)
669 group = group.strip()
671 if '-' in group:
672 """Hyphen indicates a range of values:
673 e.g. 10-20
674 """
675 items = group.split('-')
677 if len(items) == 2:
678 a = items[0]
679 b = items[1]
681 if a == b:
682 # Invalid group
683 add_error(_(f'Invalid group: {group}'))
684 continue
686 group_items = []
688 count = 0
690 a_next = a
692 while a_next is not None and a_next not in group_items:
693 group_items.append(a_next)
694 count += 1
696 # Progress to the 'next' sequential value
697 a_next = str(increment_serial_number(a_next))
699 if a_next == b:
700 # Successfully got to the end of the range
701 group_items.append(b)
702 break
704 elif count > remaining:
705 # More than the allowed number of items
706 break
708 elif a_next is None:
709 break
711 if len(group_items) > remaining:
712 add_error(
713 _(
714 f'Group range {group} exceeds allowed quantity ({expected_quantity})'
715 )
716 )
717 elif (
718 len(group_items) > 0
719 and group_items[0] == a
720 and group_items[-1] == b
721 ):
722 # In this case, the range extraction looks like it has worked
723 for item in group_items:
724 add_serial(item)
725 else:
726 add_error(_(f'Invalid group: {group}'))
728 else:
729 # In the case of a different number of hyphens, simply add the entire group
730 add_serial(group)
732 elif '+' in group:
733 """Plus character (+) indicates either:
734 - <start>+ - Expected number of serials, beginning at the specified 'start' character
735 - <start>+<num> - Specified number of serials, beginning at the specified 'start' character
736 """
737 items = group.split('+')
739 sequence_items = []
740 counter = 0
741 sequence_count = max(0, expected_quantity - len(serials))
743 if len(items) > 2 or len(items) == 0:
744 add_error(_(f'Invalid group: {group}'))
745 continue
746 elif len(items) == 2:
747 try:
748 if items[1]:
749 sequence_count = int(items[1]) + 1
750 except ValueError:
751 add_error(_(f'Invalid group: {group}'))
752 continue
754 value = items[0]
756 # Keep incrementing up to the specified quantity
757 while (
758 value is not None
759 and value not in sequence_items
760 and counter < sequence_count
761 ):
762 sequence_items.append(value)
763 value = increment_serial_number(value)
764 counter += 1
766 if len(sequence_items) == sequence_count:
767 for item in sequence_items:
768 add_serial(item)
769 else:
770 add_error(_(f'Invalid group: {group}'))
772 else:
773 # At this point, we assume that the 'group' is just a single serial value
774 add_serial(group)
776 if len(errors) > 0:
777 raise ValidationError(errors)
779 if len(serials) == 0:
780 raise ValidationError([_('No serial numbers found')])
782 if len(errors) == 0 and len(serials) != expected_quantity:
783 n = len(serials)
784 q = expected_quantity
786 raise ValidationError([
787 _(f'Number of unique serial numbers ({n}) must match quantity ({q})')
788 ])
790 return serials
793def validateFilterString(value: str, model=None) -> dict:
794 """Validate that a provided filter string looks like a list of comma-separated key=value pairs.
796 These should nominally match to a valid database filter based on the model being filtered.
798 e.g. "category=6, IPN=12"
799 e.g. "part__name=widget"
800 e.g. "item=[1,2,3], status=active"
802 The ReportTemplate class uses the filter string to work out which items a given report applies to.
803 For example, an acceptance test report template might only apply to stock items with a given IPN,
804 so the string could be set to:
806 filters = "IPN = ACME0001"
808 Returns a map of key:value pairs
809 """
810 # Empty results map
811 results = {}
813 value = str(value).strip()
815 if not value or len(value) == 0: 815 ↛ 816line 815 didn't jump to line 816 because the condition on line 815 was never true
816 return results
818 # Split by comma, but ignore commas within square brackets
819 groups = re.split(r',(?![^\[]*\])', value)
821 for group in groups: 821 ↛ 850line 821 didn't jump to line 850 because the loop on line 821 didn't complete
822 group = group.strip()
824 pair = group.split('=')
826 if len(pair) != 2: 826 ↛ 829line 826 didn't jump to line 829 because the condition on line 826 was always true
827 raise ValidationError(f'Invalid group: {group}')
829 k, v = pair
831 k = k.strip()
832 v = v.strip()
834 if not k or not v:
835 raise ValidationError(f'Invalid group: {group}')
837 # Account for 'list' support
838 if v.startswith('[') and v.endswith(']'):
839 try:
840 v = json.loads(v)
841 except json.JSONDecodeError:
842 raise ValidationError(f'Invalid list value: {v}')
844 if not isinstance(v, list):
845 raise ValidationError(f'Expected a list for key "{k}", got {type(v)}')
847 results[k] = v
849 # If a model is provided, verify that the provided filters can be used against it
850 if model is not None:
851 try:
852 model.objects.filter(**results)
853 except FieldError as e:
854 raise ValidationError(str(e))
856 return results
859def clean_decimal(number):
860 """Clean-up decimal value."""
861 # Check if empty
862 if number is None or number == '' or number == 0:
863 return Decimal(0)
865 # Convert to string and remove spaces
866 number = str(number).replace(' ', '')
868 # Guess what type of decimal and thousands separators are used
869 count_comma = number.count(',')
870 count_point = number.count('.')
872 if count_comma == 1:
873 # Comma is used as decimal separator
874 if count_point > 0:
875 # Points are used as thousands separators: remove them
876 number = number.replace('.', '')
877 # Replace decimal separator with point
878 number = number.replace(',', '.')
879 elif count_point == 1:
880 # Point is used as decimal separator
881 if count_comma > 0:
882 # Commas are used as thousands separators: remove them
883 number = number.replace(',', '')
885 # Convert to Decimal type
886 try:
887 clean_number = Decimal(number)
888 except InvalidOperation:
889 # Number cannot be converted to Decimal (eg. a string containing letters)
890 return Decimal(0)
892 return (
893 clean_number.quantize(Decimal(1))
894 if clean_number == clean_number.to_integral()
895 else clean_number.normalize()
896 )
899def strip_html_tags(value: str, raise_error=True, field_name=None):
900 """Strip HTML tags from an input string using the nh3 library.
902 If raise_error is True, a ValidationError will be thrown if HTML tags are detected
903 """
904 value = str(value).strip()
906 cleaned = nh3.clean(value, tags=frozenset())
908 # Add escaped characters back in
909 replacements = {'>': '>', '<': '<', '&': '&'}
911 for o, r in replacements.items():
912 cleaned = cleaned.replace(o, r)
914 # If the length changed, it means that HTML tags were removed!
915 if len(cleaned) != len(value) and raise_error:
916 field = field_name or 'non_field_errors'
917 raise ValidationError({field: [_('Remove HTML tags from this value')]})
919 return cleaned
922def remove_non_printable_characters(value: str, remove_newline=True) -> str:
923 """Remove non-printable / control characters from the provided string."""
924 cleaned = value
926 # Remove ASCII control characters
927 # Note that we do not sub out 0x0A (\n) here, it is done separately below
928 regex = re.compile(r'[\u0000-\u0009\u000B-\u001F\u007F-\u009F]')
929 cleaned = regex.sub('', cleaned)
931 # Remove Unicode control characters
932 regex = re.compile(r'[\u200E\u200F\u202A-\u202E]')
933 cleaned = regex.sub('', cleaned)
935 if remove_newline:
936 regex = re.compile(r'[\x0A]')
937 cleaned = regex.sub('', cleaned)
939 return cleaned
942def clean_markdown(value: str) -> str:
943 """Clean a markdown string.
945 This function will remove javascript and other potentially harmful content from the markdown string.
946 """
947 import markdown
949 try:
950 markdownify_settings = settings.MARKDOWNIFY['default']
951 except (AttributeError, KeyError):
952 markdownify_settings = {}
954 extensions = markdownify_settings.get('MARKDOWN_EXTENSIONS', [])
955 extension_configs = markdownify_settings.get('MARKDOWN_EXTENSION_CONFIGS', {})
957 # Generate raw HTML from provided markdown (without sanitizing)
958 # Note: The 'html' output_format is required to generate self closing tags, e.g. <tag> instead of <tag />
959 html = markdown.markdown(
960 value or '',
961 extensions=extensions,
962 extension_configs=extension_configs,
963 output_format='html',
964 )
966 # nh3 sanitizer settings
967 whitelist_tags = markdownify_settings.get('WHITELIST_TAGS', DEFAULT_TAGS)
968 whitelist_attrs = markdownify_settings.get('WHITELIST_ATTRS', DEAFAULT_ATTRS)
969 whitelist_styles = markdownify_settings.get('WHITELIST_STYLES', DEFAULT_CSS)
970 whitelist_protocols = markdownify_settings.get(
971 'WHITELIST_PROTOCOLS', DEFAULT_PROTOCOLS
972 )
974 # Convert bleach-style attributes (list or dict) to nh3-compatible dict format
975 if isinstance(whitelist_attrs, (list, tuple, set, frozenset)): 975 ↛ 977line 975 didn't jump to line 977 because the condition on line 975 was always true
976 attrs_dict = {'*': set(whitelist_attrs)}
977 elif isinstance(whitelist_attrs, dict):
978 attrs_dict = {tag: set(allowed) for tag, allowed in whitelist_attrs.items()}
979 else:
980 attrs_dict = None
982 # Clean the HTML content (for comparison). This must be the same as the original content
983 clean_html = nh3.clean(
984 html,
985 tags=set(whitelist_tags),
986 attributes=attrs_dict,
987 url_schemes=set(whitelist_protocols),
988 filter_style_properties=set(whitelist_styles),
989 link_rel=None,
990 strip_comments=True,
991 )
993 if html != clean_html: 993 ↛ 994line 993 didn't jump to line 994 because the condition on line 993 was never true
994 raise ValidationError(_('Data contains prohibited markdown content'))
996 return value
999def hash_barcode(barcode_data: str) -> str:
1000 """Calculate a 'unique' hash for a barcode string.
1002 This hash is used for comparison / lookup.
1004 We first remove any non-printable characters from the barcode data,
1005 as some browsers have issues scanning characters in.
1006 """
1007 barcode_data = str(barcode_data).strip()
1008 barcode_data = remove_non_printable_characters(barcode_data)
1010 barcode_hash = hashlib.md5(str(barcode_data).encode())
1012 return str(barcode_hash.hexdigest())
1015def current_time(local=True):
1016 """Return the current date and time as a datetime object.
1018 - If timezone support is active, returns a timezone aware time
1019 - If timezone support is not active, returns a timezone naive time
1021 Arguments:
1022 local: Return the time in the local timezone, otherwise UTC (default = True)
1024 """
1025 if settings.USE_TZ: 1025 ↛ 1030line 1025 didn't jump to line 1030 because the condition on line 1025 was always true
1026 now = timezone.now()
1027 now = to_local_time(now, target_tz_str=server_timezone() if local else 'UTC')
1028 return now
1029 else:
1030 return datetime.datetime.now()
1033def current_date(local=True):
1034 """Return the current date."""
1035 return current_time(local=local).date()
1038def server_timezone() -> str:
1039 """Return the timezone of the server as a string.
1041 e.g. "UTC" / "Australia/Sydney" etc
1042 """
1043 return settings.TIME_ZONE
1046def to_local_time(time, target_tz_str: Optional[str] = None):
1047 """Convert the provided time object to the local timezone.
1049 Arguments:
1050 time: The time / date to convert
1051 target_tz_str: The desired timezone (string) - defaults to server time
1053 Returns:
1054 A timezone aware datetime object, with the desired timezone
1056 Raises:
1057 TypeError: If the provided time object is not a datetime or date object
1058 """
1059 if isinstance(time, datetime.datetime): 1059 ↛ 1061line 1059 didn't jump to line 1061 because the condition on line 1059 was always true
1060 pass
1061 elif isinstance(time, datetime.date):
1062 time = timezone.datetime(year=time.year, month=time.month, day=time.day)
1063 else:
1064 raise TypeError(
1065 f'Argument must be a datetime or date object (found {type(time)}'
1066 )
1068 # Extract timezone information from the provided time
1069 source_tz = getattr(time, 'tzinfo', None)
1071 if not source_tz: 1071 ↛ 1073line 1071 didn't jump to line 1073 because the condition on line 1071 was never true
1072 # Default to UTC if not provided
1073 source_tz = ZoneInfo('UTC')
1075 if not target_tz_str: 1075 ↛ 1076line 1075 didn't jump to line 1076 because the condition on line 1075 was never true
1076 target_tz_str = server_timezone()
1078 try:
1079 target_tz = ZoneInfo(str(target_tz_str))
1080 except ZoneInfoNotFoundError:
1081 target_tz = ZoneInfo('UTC')
1083 target_time = time.replace(tzinfo=source_tz).astimezone(target_tz)
1085 return target_time
1088def get_objectreference(
1089 obj, type_ref: str = 'content_type', object_ref: str = 'object_id'
1090):
1091 """Lookup method for the GenericForeignKey fields.
1093 Attributes:
1094 - obj: object that will be resolved
1095 - type_ref: field name for the contenttype field in the model
1096 - object_ref: field name for the object id in the model
1098 Example implementation in the serializer:
1099 ```
1100 target = serializers.SerializerMethodField()
1101 def get_target(self, obj):
1102 return get_objectreference(obj, 'target_content_type', 'target_object_id')
1103 ```
1105 The method name must always be the name of the field prefixed by 'get_'
1106 """
1107 model_cls = getattr(obj, type_ref)
1108 obj_id = getattr(obj, object_ref)
1110 # check if references are set -> return nothing if not
1111 if model_cls is None or obj_id is None: 1111 ↛ 1112line 1111 didn't jump to line 1112 because the condition on line 1111 was never true
1112 return None
1114 # resolve referenced data into objects
1115 model_cls = model_cls.model_class()
1117 try:
1118 item = model_cls.objects.get(id=obj_id)
1119 except model_cls.DoesNotExist:
1120 return None
1122 url_fnc = getattr(item, 'get_absolute_url', None)
1124 # create output
1125 ret = {}
1126 if url_fnc: 1126 ↛ 1127line 1126 didn't jump to line 1127 because the condition on line 1126 was never true
1127 ret['link'] = url_fnc()
1129 return {
1130 'name': str(item),
1131 'model_name': str(model_cls._meta.verbose_name),
1132 'model_type': str(model_cls._meta.model_name),
1133 'model_id': getattr(item, 'pk', None),
1134 **ret,
1135 }
1138Inheritors_T = TypeVar('Inheritors_T')
1141def inheritors(
1142 cls: type[Inheritors_T], subclasses: bool = True
1143) -> set[type[Inheritors_T]]:
1144 """Return all classes that are subclasses from the supplied cls.
1146 Args:
1147 cls: The class to search for subclasses
1148 subclasses: Include subclasses of subclasses (default = True)
1149 """
1150 subcls = set()
1151 work = [cls]
1153 while work:
1154 parent = work.pop()
1155 for child in parent.__subclasses__():
1156 if child not in subcls: 1156 ↛ 1155line 1156 didn't jump to line 1155 because the condition on line 1156 was always true
1157 subcls.add(child)
1158 if subclasses: 1158 ↛ 1155line 1158 didn't jump to line 1155 because the condition on line 1158 was always true
1159 work.append(child)
1160 return subcls
1163def pui_url(subpath: str) -> str:
1164 """Return the URL for a web subpath."""
1165 if not subpath.startswith('/'): 1165 ↛ 1166line 1165 didn't jump to line 1166 because the condition on line 1165 was never true
1166 subpath = '/' + subpath
1167 return f'/{settings.FRONTEND_URL_BASE}{subpath}'
1170def plugins_info(*args, **kwargs):
1171 """Return information about activated plugins."""
1172 from plugin import PluginMixinEnum
1173 from plugin.registry import registry
1175 # Check if plugins are even enabled
1176 if not settings.PLUGINS_ENABLED: 1176 ↛ 1180line 1176 didn't jump to line 1180 because the condition on line 1176 was always true
1177 return False
1179 # Fetch active plugins
1180 plugins = registry.with_mixin(PluginMixinEnum.BASE)
1182 # Format list
1183 return [
1184 {'name': plg.name, 'slug': plg.slug, 'version': plg.version} for plg in plugins
1185 ]
1188def sanitize_token(token_value: str, front=8, back=12) -> str:
1189 """Sanitize a token by replacing the middle characters with asterisks.
1191 Args:
1192 token_value: The token string to sanitize
1193 front: Number of characters to show at the start of the token (default = 8)
1194 back: Number of characters to show at the end of the token (default = 12)
1196 Returns:
1197 The sanitized token string
1198 """
1199 middle = len(token_value) - (front + back)
1200 return token_value[:front] + '*' * middle + token_value[-back:]