Coverage for src/backend/InvenTree/InvenTree/helpers_model.py: 38%

133 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 that access the database.""" 

2 

3import io 

4import ipaddress 

5import socket 

6from typing import Optional, cast 

7from urllib.parse import urljoin, urlparse 

8 

9from django.conf import settings 

10from django.core.validators import URLValidator 

11from django.db.utils import OperationalError, ProgrammingError 

12from django.utils.translation import gettext_lazy as _ 

13 

14import requests 

15import requests.exceptions 

16import structlog 

17from PIL import Image 

18 

19from common.notifications import ( 

20 InvenTreeNotificationBodies, 

21 NotificationBody, 

22 trigger_notification, 

23) 

24from common.settings import get_global_setting 

25from InvenTree.cache import ( 

26 get_cached_content_types, 

27 get_session_cache, 

28 set_session_cache, 

29) 

30from InvenTree.ready import ignore_ready_warning 

31 

32logger = structlog.get_logger('inventree') 

33 

34 

35def get_base_url(request=None) -> str: 

36 """Return the base URL for the InvenTree server. 

37 

38 The base URL is determined in the following order of decreasing priority: 

39 

40 1. If a request object is provided, use the request URL 

41 2. Multi-site is enabled, and the current site has a valid URL 

42 3. If settings.SITE_URL is set (e.g. in the Django settings), use that 

43 4. If the InvenTree setting INVENTREE_BASE_URL is set, use that 

44 """ 

45 # Check if a request is provided 

46 if request: 46 ↛ 47line 46 didn't jump to line 47 because the condition on line 46 was never true

47 return request.build_absolute_uri('/') 

48 

49 # Check if multi-site is enabled 

50 try: 

51 from django.contrib.sites.models import Site 

52 

53 return Site.objects.get_current().domain 

54 except (ImportError, RuntimeError): 

55 pass 

56 

57 # Check if a global site URL is provided 

58 if site_url := getattr(settings, 'SITE_URL', None): 58 ↛ 62line 58 didn't jump to line 62 because the condition on line 58 was always true

59 return site_url 

60 

61 # Check if a global InvenTree setting is provided 

62 try: 

63 if site_url := get_global_setting('INVENTREE_BASE_URL', create=False): 

64 return cast(str, site_url) 

65 except (ProgrammingError, OperationalError): 

66 pass 

67 

68 # No base URL available 

69 return '' 

70 

71 

72def construct_absolute_url(*arg, base_url=None, request=None): 

73 """Construct (or attempt to construct) an absolute URL from a relative URL. 

74 

75 Args: 

76 *arg: The relative URL to construct 

77 base_url: The base URL to use for the construction (if not provided, will attempt to determine from settings) 

78 request: The request object to use for the construction (optional) 

79 """ 

80 relative_url = '/'.join(arg) 

81 

82 if not base_url: 82 ↛ 85line 82 didn't jump to line 85 because the condition on line 82 was always true

83 base_url = get_base_url(request=request) 

84 

85 return urljoin(base_url, relative_url) 

86 

87 

88def validate_url_no_ssrf(url): 

89 """Validate that a URL does not point to a private/internal network address. 

90 

91 Resolves the hostname to an IP address and checks it against private, 

92 loopback, link-local, and reserved IP ranges to prevent SSRF attacks. 

93 

94 Arguments: 

95 url: The URL to validate 

96 

97 Raises: 

98 ValueError: If the URL resolves to a private or reserved IP address 

99 """ 

100 parsed = urlparse(url) 

101 hostname = parsed.hostname 

102 

103 if not hostname: 

104 raise ValueError(_('Invalid URL: no hostname')) 

105 

106 try: 

107 addrinfo = socket.getaddrinfo(hostname, None) 

108 except socket.gaierror: 

109 raise ValueError(_('Invalid URL: hostname could not be resolved')) 

110 

111 for _family, _type, _proto, _canonname, sockaddr in addrinfo: 

112 ip = ipaddress.ip_address(sockaddr[0]) 

113 

114 if ip.is_private or ip.is_loopback or ip.is_link_local or ip.is_reserved: 

115 raise ValueError(_('URL points to a private or reserved IP address')) 

116 

117 

118def download_image_from_url( 

119 remote_url: str, 

120 timeout: float = 2.5, 

121 user_agent: str = '', 

122 max_size: Optional[int] = None, 

123): 

124 """Download an image file from a remote URL. 

125 

126 This is a potentially dangerous operation, so we must perform some checks: 

127 - The remote URL is available 

128 - The Content-Length is provided, and is not too large 

129 - The file is a valid image file 

130 

131 Arguments: 

132 remote_url: The remote URL to retrieve image 

133 timeout: Connection timeout in seconds (default = 5) 

134 user_agent: User-Agent string to use for the request (optional) 

135 max_size: Maximum allowed image size (in bytes) (default = 1MB) 

136 

137 Returns: 

138 An in-memory PIL image file, if the download was successful 

139 

140 Raises: 

141 requests.exceptions.ConnectionError: Connection could not be established 

142 requests.exceptions.Timeout: Connection timed out 

143 requests.exceptions.HTTPError: Server responded with invalid response code 

144 ValueError: Server responded with invalid 'Content-Length' value 

145 TypeError: Response is not a valid image 

146 """ 

147 # Check that the provided URL at least looks valid 

148 validator = URLValidator() 

149 validator(remote_url) 

150 

151 # SSRF protection: validate the resolved IP is not private/internal 

152 validate_url_no_ssrf(remote_url) 

153 

154 # Calculate maximum allowable image size (in bytes) 

155 max_size = max_size or 1 * 1024 * 1024 # Default to 1MB if not provided 

156 

157 # Add user specified user-agent to request (if specified) 

158 headers = {'User-Agent': user_agent} if user_agent else None 

159 

160 try: 

161 response = requests.get( 

162 remote_url, 

163 timeout=timeout, 

164 allow_redirects=False, 

165 stream=True, 

166 headers=headers, 

167 ) 

168 

169 # Handle redirects manually to validate each destination 

170 max_redirects = 5 

171 redirect_count = 0 

172 

173 while response.is_redirect and redirect_count < max_redirects: 

174 redirect_url = response.headers.get('Location') 

175 if not redirect_url: 

176 break 

177 

178 # Validate the redirect destination against SSRF 

179 validator(redirect_url) 

180 validate_url_no_ssrf(redirect_url) 

181 

182 redirect_count += 1 

183 response = requests.get( 

184 redirect_url, 

185 timeout=timeout, 

186 allow_redirects=False, 

187 stream=True, 

188 headers=headers, 

189 ) 

190 

191 if redirect_count >= max_redirects: 

192 raise ValueError(_('Too many redirects')) 

193 

194 # Throw an error if anything goes wrong 

195 response.raise_for_status() 

196 except requests.exceptions.ConnectionError as exc: 

197 raise Exception(_('Connection error') + f': {exc!s}') 

198 except requests.exceptions.Timeout as exc: 

199 raise exc 

200 except requests.exceptions.HTTPError: 

201 raise requests.exceptions.HTTPError( 

202 _('Server responded with invalid status code') + f': {response.status_code}' 

203 ) 

204 except ValueError: 

205 raise 

206 except Exception as exc: 

207 raise Exception(_('Exception occurred') + f': {exc!s}') 

208 

209 if response.status_code != 200: 

210 raise Exception( 

211 _('Server responded with invalid status code') + f': {response.status_code}' 

212 ) 

213 

214 try: 

215 content_length = int(response.headers.get('Content-Length', 0)) 

216 except ValueError: 

217 raise ValueError(_('Server responded with invalid Content-Length value')) 

218 

219 if content_length > max_size: 

220 raise ValueError(_('Image size is too large')) 

221 

222 # Download the file, ensuring we do not exceed the reported size 

223 file = io.BytesIO() 

224 

225 dl_size = 0 

226 chunk_size = 64 * 1024 

227 

228 for chunk in response.iter_content(chunk_size=chunk_size): 

229 dl_size += len(chunk) 

230 

231 if dl_size > max_size: 

232 raise ValueError(_('Image download exceeded maximum size')) 

233 

234 file.write(chunk) 

235 

236 if dl_size == 0: 

237 raise ValueError(_('Remote server returned empty response')) 

238 

239 # Now, attempt to convert the downloaded data to a valid image file 

240 # img.verify() will throw an exception if the image is not valid 

241 try: 

242 img = Image.open(file).convert() 

243 img.verify() 

244 except Exception: 

245 raise TypeError(_('Supplied URL is not a valid image file')) 

246 

247 return img 

248 

249 

250@ignore_ready_warning 

251def getModelsWithMixin(mixin_class) -> list: 

252 """Return a list of database models that inherit from the given mixin class. 

253 

254 Args: 

255 mixin_class: The mixin class to search for 

256 Returns: 

257 List of models that inherit from the given mixin class 

258 """ 

259 # First, look in the session cache - to prevent repeated expensive comparisons 

260 cache_key = f'models_with_mixin_{mixin_class.__name__}' 

261 

262 if cached_models := get_session_cache(cache_key): 

263 return cached_models 

264 

265 content_types = get_cached_content_types() 

266 

267 db_models = [x.model_class() for x in content_types if x is not None] 

268 

269 models_with_mixin = [ 

270 x for x in db_models if x is not None and issubclass(x, mixin_class) 

271 ] 

272 # sort to make resulting list deterministic (and easier to test) 

273 models_with_mixin.sort(key=lambda x: x._meta.label_lower) 

274 

275 # Store the result in the session cache 

276 set_session_cache(cache_key, models_with_mixin) 

277 return models_with_mixin 

278 

279 

280def notify_responsible( 

281 instance, 

282 sender, 

283 content: NotificationBody = InvenTreeNotificationBodies.NewOrder, 

284 exclude=None, 

285 extra_users: Optional[list] = None, 

286): 

287 """Notify all responsible parties of a change in an instance. 

288 

289 Parses the supplied content with the provided instance and sender and sends a notification to all responsible users, 

290 excluding the optional excluded list. 

291 

292 Args: 

293 instance: The newly created instance 

294 sender: Sender model reference 

295 content (NotificationBody, optional): _description_. Defaults to InvenTreeNotificationBodies.NewOrder. 

296 exclude (User, optional): User instance that should be excluded. Defaults to None. 

297 extra_users (list, optional): List of extra users to notify. Defaults to None. 

298 """ 

299 import InvenTree.ready 

300 

301 if InvenTree.ready.isImportingData() or InvenTree.ready.isRunningMigrations(): 301 ↛ 302line 301 didn't jump to line 302 because the condition on line 301 was never true

302 return 

303 

304 users = [instance.responsible] 

305 

306 if extra_users: 306 ↛ 307line 306 didn't jump to line 307 because the condition on line 306 was never true

307 users.extend(extra_users) 

308 

309 notify_users(users, instance, sender, content=content, exclude=exclude) 

310 

311 

312def notify_users( 

313 users, 

314 instance, 

315 sender, 

316 content: NotificationBody = InvenTreeNotificationBodies.NewOrder, 

317 exclude=None, 

318): 

319 """Notify all passed users or groups. 

320 

321 Parses the supplied content with the provided instance and sender and sends a notification to all users, 

322 excluding the optional excluded list. 

323 

324 Args: 

325 users: List of users or groups to notify 

326 instance: The newly created instance 

327 sender: Sender model reference 

328 content (NotificationBody, optional): _description_. Defaults to InvenTreeNotificationBodies.NewOrder. 

329 exclude (User, optional): User instance that should be excluded. Defaults to None. 

330 """ 

331 # Setup context for notification parsing 

332 content_context = { 

333 'instance': str(instance), 

334 'verbose_name': sender._meta.verbose_name, 

335 'app_label': sender._meta.app_label, 

336 'model_name': sender._meta.model_name, 

337 } 

338 

339 # Setup notification context 

340 context = { 

341 'instance': instance, 

342 'name': content.name.format(**content_context), 

343 'message': content.message.format(**content_context), 

344 'link': construct_absolute_url(instance.get_absolute_url()), 

345 'template': {'subject': content.name.format(**content_context)}, 

346 } 

347 

348 tmp = content.template 

349 if tmp: 349 ↛ 353line 349 didn't jump to line 353 because the condition on line 349 was always true

350 context['template']['html'] = tmp.format(**content_context) 

351 

352 # Create notification 

353 trigger_notification( 

354 instance, 

355 content.slug.format(**content_context), 

356 targets=users, 

357 target_exclude=[exclude], 

358 context=context, 

359 )