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

1"""Provides helper functions used throughout the InvenTree project.""" 

2 

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 

15 

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 _ 

24 

25import nh3 

26import structlog 

27from djmoney.money import Money 

28from PIL import Image 

29from stdimage.models import StdImageField, StdImageFieldFile 

30 

31from common.currency import currency_code_default 

32from InvenTree.sanitizer import ( 

33 DEAFAULT_ATTRS, 

34 DEFAULT_CSS, 

35 DEFAULT_PROTOCOLS, 

36 DEFAULT_TAGS, 

37) 

38 

39logger = structlog.get_logger('inventree') 

40 

41INT_CLIP_MAX = 0x7FFFFFFF 

42 

43 

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. 

48 

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 

57 

58 def do_clip(value: int, clip: int, allow_negative: bool) -> int: 

59 """Perform clipping on the provided value. 

60 

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 

68 

69 clip = min(clip, INT_CLIP_MAX) 

70 

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 

75 

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) 

78 

79 return value 

80 

81 reference = str(reference).strip() 

82 

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 

86 

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 

93 

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 

101 

102 # Look at the start of the string - can it be "integerized"? 

103 result = re.match(r'^(\d+)', reference) 

104 

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) 

114 

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 

121 

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) 

125 

126 if not allow_negative and ref_int < 0: 

127 ref_int = abs(ref_int) 

128 

129 return ref_int 

130 

131 

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. 

134 

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 = '' 

139 

140 key = test_name.strip().lower() 

141 key = key.replace(' ', '') 

142 

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 

147 

148 if char.isidentifier(): 

149 return True 

150 

151 return bool(char.isalnum()) 

152 

153 # Remove any characters that cannot be used to represent a variable 

154 key = ''.join([c for c in key if valid_char(c)]) 

155 

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 

159 

160 return key 

161 

162 

163def constructPathString(path: list[str], max_chars: int = 250) -> str: 

164 """Construct a 'path string' for the given path. 

165 

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) 

171 

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:] 

176 

177 return pathstring 

178 

179 

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) 

190 

191 return default_storage.url(file.name) 

192 

193 

194def regenerate_imagefile(_file, _name: str): 

195 """Regenerate a file object for a given variation name. 

196 

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] 

203 

204 

205def image2name(img_obj: StdImageField, do_preview: bool, do_thumbnail: bool): 

206 """Convert an image object to a filename string. 

207 

208 Arguments: 

209 img_obj: Image object 

210 do_preview: Return preview image name 

211 do_thumbnail: Return thumbnail image name 

212 """ 

213 

214 def extract(ref: str): 

215 return None if not hasattr(img_obj, ref) else getattr(img_obj, ref).name 

216 

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 

225 

226 

227def getStaticUrl(filename): 

228 """Return the qualified access path for the given file, under the static media directory.""" 

229 return StaticFilesStorage().url(filename) 

230 

231 

232def TestIfImage(img) -> bool: 

233 """Test if an image file is indeed an image. 

234 

235 Arguments: 

236 img: A file-like object 

237 

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 

246 

247 

248def getBlankImage(): 

249 """Return the qualified path for the 'blank image' placeholder.""" 

250 return getStaticUrl('img/blank_image.png') 

251 

252 

253def getBlankThumbnail(): 

254 """Return the qualified path for the 'blank image' thumbnail placeholder.""" 

255 return getStaticUrl('img/blank_image.thumbnail.png') 

256 

257 

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)) 

263 

264 

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() 

269 

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 

276 

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) 

281 

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') 

287 

288 

289def getSplashScreen(custom=True): 

290 """Return the InvenTree splash screen, or a custom splash if available.""" 

291 static_storage = StaticFilesStorage() 

292 

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) 

296 

297 # No custom splash screen 

298 return static_storage.url('img/inventree_splash.jpg') 

299 

300 

301def getCustomOption(reference: str): 

302 """Return the value of a custom option from settings.CUSTOMIZE. 

303 

304 Args: 

305 reference: Reference key for the custom option 

306 """ 

307 return settings.CUSTOMIZE.get(reference, None) 

308 

309 

310def TestIfImageURL(url): 

311 """Test if an image URL (or filename) looks like a valid image format. 

312 

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 ] 

326 

327 

328def str2bool(text, test=True) -> bool: 

329 """Test if a string 'looks' like a boolean value. 

330 

331 Args: 

332 text: Input text 

333 test (default = True): Set which boolean value to look for 

334 

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'] 

341 

342 

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) 

346 

347 

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. 

350 

351 Args: 

352 text: Input text 

353 

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 ] 

366 

367 

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) 

372 

373 if rounding is not None: 

374 d = round(d, rounding) 

375 

376 d = d.normalize() 

377 

378 # Ref: https://docs.python.org/3/library/decimal.html 

379 return d.quantize(Decimal(1)) if d == d.to_integral() else d.normalize() 

380 

381 

382def increment(value): 

383 """Attempt to increment an integer (or a string that looks like an integer). 

384 

385 e.g. 

386 

387 001 -> 002 

388 2 -> 3 

389 AB01 -> AB02 

390 QQQ -> QQQ 

391 

392 """ 

393 # Ignore empty strings 

394 if value in ['', None]: 

395 # Provide a default value if provided with a null input 

396 return '1' 

397 

398 value = str(value).strip() 

399 

400 pattern = r'(.*?)(\d+)?$' 

401 

402 result = re.search(pattern, value) 

403 

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 

407 

408 groups = result.groups() 

409 

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 

413 

414 prefix, number = groups 

415 

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 

419 

420 # Record the width of the number 

421 width = len(number) 

422 

423 try: 

424 number = int(number) + 1 

425 number = str(number) 

426 except ValueError: 

427 pass 

428 

429 return prefix + str(number).zfill(width) 

430 

431 

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. 

434 

435 Args: 

436 d: A python Decimal object 

437 

438 Returns: 

439 A string representation of the input number 

440 """ 

441 if type(d) is Decimal: 

442 d = normalize(d) 

443 

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) 

450 

451 s = str(d) 

452 

453 # Return entire number if there is no decimal place 

454 if '.' not in s: 

455 return s 

456 

457 return s.rstrip('0').rstrip('.') 

458 

459 

460def decimal2money(d, currency=None): 

461 """Format a Decimal number as Money. 

462 

463 Args: 

464 d: A python Decimal object 

465 currency: Currency of the input amount, defaults to default currency in settings 

466 

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) 

473 

474 

475def WrapWithQuotes(text, quote='"'): 

476 """Wrap the supplied text with quotes. 

477 

478 Args: 

479 text: Input text to wrap 

480 quote: Quote character to use for wrapping (default = "") 

481 

482 Returns: 

483 Supplied text wrapped in quote char 

484 """ 

485 if not text.startswith(quote): 

486 text = quote + text 

487 

488 if not text.endswith(quote): 

489 text = text + quote 

490 

491 return text 

492 

493 

494def GetExportOptions() -> list: 

495 """Return a set of allowable import / export file formats.""" 

496 return [['csv', 'CSV'], ['xlsx', 'Excel'], ['tsv', 'TSV']] 

497 

498 

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()] 

502 

503 

504def DownloadFile( 

505 data, filename, content_type='application/text', inline=False 

506) -> StreamingHttpResponse: 

507 """Create a dynamic file for the user to download. 

508 

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) 

514 

515 Return: 

516 A StreamingHttpResponse object wrapping the supplied data 

517 """ 

518 filename = WrapWithQuotes(filename) 

519 length = len(data) 

520 

521 if isinstance(data, str): 

522 wrapper = FileWrapper(io.StringIO(data)) 

523 else: 

524 wrapper = FileWrapper(io.BytesIO(data)) 

525 

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 

530 

531 if inline: 

532 disposition = f'inline; filename={filename}' 

533 else: 

534 disposition = f'attachment; filename={filename}' 

535 

536 response['Content-Disposition'] = disposition 

537 return response 

538 

539 

540def increment_serial_number(serial, part=None): 

541 """Given a serial number, (attempt to) generate the *next* serial number. 

542 

543 Note: This method is exposed to custom plugins. 

544 

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 

548 

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 

555 

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() 

559 

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 

566 

567 signature = inspect.signature(plugin.increment_serial_number) 

568 

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) 

578 

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) 

582 

583 

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. 

588 

589 The input string can be specified using the following concepts: 

590 

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> 

596 

597 Actual generation of sequential serials is passed to the 'validation' plugin mixin, 

598 allowing custom plugins to determine how serial values are incremented. 

599 

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) 

608 

609 try: 

610 expected_quantity = int(expected_quantity) 

611 except ValueError: 

612 raise ValidationError([_('Invalid quantity provided')]) 

613 

614 if expected_quantity > 1000: 

615 raise ValidationError({ 

616 'quantity': [_('Cannot serialize more than 1000 items at once')] 

617 }) 

618 

619 input_string = str(input_string).strip() if input_string else '' 

620 

621 if len(input_string) == 0: 

622 raise ValidationError([_('Empty serial number string')]) 

623 

624 next_value = increment_serial_number(starting_value, part=part) 

625 

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) 

630 

631 # Split input string by whitespace or comma (,) characters 

632 groups = re.split(r'[\s,]+', input_string) 

633 

634 serials = [] 

635 errors = [] 

636 

637 def add_error(error: str): 

638 """Helper function for adding an error message.""" 

639 if error not in errors: 

640 errors.append(error) 

641 

642 def add_serial(serial): 

643 """Helper function to check for duplicated values.""" 

644 serial = serial.strip() 

645 

646 # Ignore blank / empty serials 

647 if len(serial) == 0: 

648 return 

649 

650 if serial in serials: 

651 add_error(_('Duplicate serial') + f': {serial}') 

652 else: 

653 serials.append(serial) 

654 

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) 

659 

660 if len(errors) > 0: 

661 raise ValidationError(errors) 

662 else: 

663 return serials 

664 

665 for group in groups: 

666 # Calculate the "remaining" quantity of serial numbers 

667 remaining = expected_quantity - len(serials) 

668 

669 group = group.strip() 

670 

671 if '-' in group: 

672 """Hyphen indicates a range of values: 

673 e.g. 10-20 

674 """ 

675 items = group.split('-') 

676 

677 if len(items) == 2: 

678 a = items[0] 

679 b = items[1] 

680 

681 if a == b: 

682 # Invalid group 

683 add_error(_(f'Invalid group: {group}')) 

684 continue 

685 

686 group_items = [] 

687 

688 count = 0 

689 

690 a_next = a 

691 

692 while a_next is not None and a_next not in group_items: 

693 group_items.append(a_next) 

694 count += 1 

695 

696 # Progress to the 'next' sequential value 

697 a_next = str(increment_serial_number(a_next)) 

698 

699 if a_next == b: 

700 # Successfully got to the end of the range 

701 group_items.append(b) 

702 break 

703 

704 elif count > remaining: 

705 # More than the allowed number of items 

706 break 

707 

708 elif a_next is None: 

709 break 

710 

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}')) 

727 

728 else: 

729 # In the case of a different number of hyphens, simply add the entire group 

730 add_serial(group) 

731 

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('+') 

738 

739 sequence_items = [] 

740 counter = 0 

741 sequence_count = max(0, expected_quantity - len(serials)) 

742 

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 

753 

754 value = items[0] 

755 

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 

765 

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}')) 

771 

772 else: 

773 # At this point, we assume that the 'group' is just a single serial value 

774 add_serial(group) 

775 

776 if len(errors) > 0: 

777 raise ValidationError(errors) 

778 

779 if len(serials) == 0: 

780 raise ValidationError([_('No serial numbers found')]) 

781 

782 if len(errors) == 0 and len(serials) != expected_quantity: 

783 n = len(serials) 

784 q = expected_quantity 

785 

786 raise ValidationError([ 

787 _(f'Number of unique serial numbers ({n}) must match quantity ({q})') 

788 ]) 

789 

790 return serials 

791 

792 

793def validateFilterString(value: str, model=None) -> dict: 

794 """Validate that a provided filter string looks like a list of comma-separated key=value pairs. 

795 

796 These should nominally match to a valid database filter based on the model being filtered. 

797 

798 e.g. "category=6, IPN=12" 

799 e.g. "part__name=widget" 

800 e.g. "item=[1,2,3], status=active" 

801 

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: 

805 

806 filters = "IPN = ACME0001" 

807 

808 Returns a map of key:value pairs 

809 """ 

810 # Empty results map 

811 results = {} 

812 

813 value = str(value).strip() 

814 

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 

817 

818 # Split by comma, but ignore commas within square brackets 

819 groups = re.split(r',(?![^\[]*\])', value) 

820 

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() 

823 

824 pair = group.split('=') 

825 

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}') 

828 

829 k, v = pair 

830 

831 k = k.strip() 

832 v = v.strip() 

833 

834 if not k or not v: 

835 raise ValidationError(f'Invalid group: {group}') 

836 

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}') 

843 

844 if not isinstance(v, list): 

845 raise ValidationError(f'Expected a list for key "{k}", got {type(v)}') 

846 

847 results[k] = v 

848 

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)) 

855 

856 return results 

857 

858 

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) 

864 

865 # Convert to string and remove spaces 

866 number = str(number).replace(' ', '') 

867 

868 # Guess what type of decimal and thousands separators are used 

869 count_comma = number.count(',') 

870 count_point = number.count('.') 

871 

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(',', '') 

884 

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) 

891 

892 return ( 

893 clean_number.quantize(Decimal(1)) 

894 if clean_number == clean_number.to_integral() 

895 else clean_number.normalize() 

896 ) 

897 

898 

899def strip_html_tags(value: str, raise_error=True, field_name=None): 

900 """Strip HTML tags from an input string using the nh3 library. 

901 

902 If raise_error is True, a ValidationError will be thrown if HTML tags are detected 

903 """ 

904 value = str(value).strip() 

905 

906 cleaned = nh3.clean(value, tags=frozenset()) 

907 

908 # Add escaped characters back in 

909 replacements = {'&gt;': '>', '&lt;': '<', '&amp;': '&'} 

910 

911 for o, r in replacements.items(): 

912 cleaned = cleaned.replace(o, r) 

913 

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')]}) 

918 

919 return cleaned 

920 

921 

922def remove_non_printable_characters(value: str, remove_newline=True) -> str: 

923 """Remove non-printable / control characters from the provided string.""" 

924 cleaned = value 

925 

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) 

930 

931 # Remove Unicode control characters 

932 regex = re.compile(r'[\u200E\u200F\u202A-\u202E]') 

933 cleaned = regex.sub('', cleaned) 

934 

935 if remove_newline: 

936 regex = re.compile(r'[\x0A]') 

937 cleaned = regex.sub('', cleaned) 

938 

939 return cleaned 

940 

941 

942def clean_markdown(value: str) -> str: 

943 """Clean a markdown string. 

944 

945 This function will remove javascript and other potentially harmful content from the markdown string. 

946 """ 

947 import markdown 

948 

949 try: 

950 markdownify_settings = settings.MARKDOWNIFY['default'] 

951 except (AttributeError, KeyError): 

952 markdownify_settings = {} 

953 

954 extensions = markdownify_settings.get('MARKDOWN_EXTENSIONS', []) 

955 extension_configs = markdownify_settings.get('MARKDOWN_EXTENSION_CONFIGS', {}) 

956 

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 ) 

965 

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 ) 

973 

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 

981 

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 ) 

992 

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')) 

995 

996 return value 

997 

998 

999def hash_barcode(barcode_data: str) -> str: 

1000 """Calculate a 'unique' hash for a barcode string. 

1001 

1002 This hash is used for comparison / lookup. 

1003 

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) 

1009 

1010 barcode_hash = hashlib.md5(str(barcode_data).encode()) 

1011 

1012 return str(barcode_hash.hexdigest()) 

1013 

1014 

1015def current_time(local=True): 

1016 """Return the current date and time as a datetime object. 

1017 

1018 - If timezone support is active, returns a timezone aware time 

1019 - If timezone support is not active, returns a timezone naive time 

1020 

1021 Arguments: 

1022 local: Return the time in the local timezone, otherwise UTC (default = True) 

1023 

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() 

1031 

1032 

1033def current_date(local=True): 

1034 """Return the current date.""" 

1035 return current_time(local=local).date() 

1036 

1037 

1038def server_timezone() -> str: 

1039 """Return the timezone of the server as a string. 

1040 

1041 e.g. "UTC" / "Australia/Sydney" etc 

1042 """ 

1043 return settings.TIME_ZONE 

1044 

1045 

1046def to_local_time(time, target_tz_str: Optional[str] = None): 

1047 """Convert the provided time object to the local timezone. 

1048 

1049 Arguments: 

1050 time: The time / date to convert 

1051 target_tz_str: The desired timezone (string) - defaults to server time 

1052 

1053 Returns: 

1054 A timezone aware datetime object, with the desired timezone 

1055 

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 ) 

1067 

1068 # Extract timezone information from the provided time 

1069 source_tz = getattr(time, 'tzinfo', None) 

1070 

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') 

1074 

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() 

1077 

1078 try: 

1079 target_tz = ZoneInfo(str(target_tz_str)) 

1080 except ZoneInfoNotFoundError: 

1081 target_tz = ZoneInfo('UTC') 

1082 

1083 target_time = time.replace(tzinfo=source_tz).astimezone(target_tz) 

1084 

1085 return target_time 

1086 

1087 

1088def get_objectreference( 

1089 obj, type_ref: str = 'content_type', object_ref: str = 'object_id' 

1090): 

1091 """Lookup method for the GenericForeignKey fields. 

1092 

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 

1097 

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 ``` 

1104 

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) 

1109 

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 

1113 

1114 # resolve referenced data into objects 

1115 model_cls = model_cls.model_class() 

1116 

1117 try: 

1118 item = model_cls.objects.get(id=obj_id) 

1119 except model_cls.DoesNotExist: 

1120 return None 

1121 

1122 url_fnc = getattr(item, 'get_absolute_url', None) 

1123 

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() 

1128 

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 } 

1136 

1137 

1138Inheritors_T = TypeVar('Inheritors_T') 

1139 

1140 

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. 

1145 

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] 

1152 

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 

1161 

1162 

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}' 

1168 

1169 

1170def plugins_info(*args, **kwargs): 

1171 """Return information about activated plugins.""" 

1172 from plugin import PluginMixinEnum 

1173 from plugin.registry import registry 

1174 

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 

1178 

1179 # Fetch active plugins 

1180 plugins = registry.with_mixin(PluginMixinEnum.BASE) 

1181 

1182 # Format list 

1183 return [ 

1184 {'name': plg.name, 'slug': plg.slug, 'version': plg.version} for plg in plugins 

1185 ] 

1186 

1187 

1188def sanitize_token(token_value: str, front=8, back=12) -> str: 

1189 """Sanitize a token by replacing the middle characters with asterisks. 

1190 

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) 

1195 

1196 Returns: 

1197 The sanitized token string 

1198 """ 

1199 middle = len(token_value) - (front + back) 

1200 return token_value[:front] + '*' * middle + token_value[-back:]