Coverage for netbox/middleware.py: 44%

152 statements  

« prev     ^ index     » next       coverage.py v7.15.2, created at 2026-10-10 18:35 +0000

1import logging 

2import uuid 

3 

4from django.conf import settings 

5from django.contrib import auth, messages 

6from django.contrib.auth.middleware import RemoteUserMiddleware as RemoteUserMiddleware_ 

7from django.core.exceptions import ImproperlyConfigured 

8from django.core.signals import got_request_exception 

9from django.db import ProgrammingError, connection 

10from django.db.utils import InternalError 

11from django.http import Http404, HttpResponseRedirect 

12from django.middleware.common import CommonMiddleware as DjangoCommonMiddleware 

13from django.utils.translation import gettext_lazy as _ 

14from django_prometheus import middleware 

15from social_django.middleware import SocialAuthExceptionMiddleware as SocialAuthExceptionMiddleware_ 

16 

17from netbox.config import clear_config, get_config 

18from netbox.metrics import Metrics 

19from netbox.views import handler_500 

20from utilities.api import is_api_request, is_graphql_request 

21from utilities.error_handlers import handle_rest_api_exception 

22from utilities.request import apply_request_processors 

23 

24__all__ = ( 

25 'CommonMiddleware', 

26 'CoreMiddleware', 

27 'MaintenanceModeMiddleware', 

28 'PrometheusAfterMiddleware', 

29 'PrometheusBeforeMiddleware', 

30 'RemoteUserMiddleware', 

31 'SocialAuthExceptionMiddleware', 

32) 

33 

34 

35class CommonMiddleware(DjangoCommonMiddleware): 

36 """ 

37 Subclass of Django's CommonMiddleware that suppresses the APPEND_SLASH 

38 redirect for REST API requests using an unsafe HTTP method. Redirecting a 

39 POST/PUT/PATCH/DELETE to a trailing-slash URL would either drop the request 

40 body (clients downgrade to GET on a 302) or raise a RuntimeError when 

41 DEBUG is enabled. Letting the original 404 propagate gives the caller a 

42 clear, actionable error instead. 

43 """ 

44 UNSAFE_METHODS = frozenset(('DELETE', 'PATCH', 'POST', 'PUT')) 

45 

46 def should_redirect_with_slash(self, request): 

47 if request.method in self.UNSAFE_METHODS and is_api_request(request): 

48 return False 

49 return super().should_redirect_with_slash(request) 

50 

51 

52class CoreMiddleware: 

53 

54 def __init__(self, get_response): 

55 self.get_response = get_response 

56 

57 def __call__(self, request): 

58 

59 # Assign a random unique ID to the request. This will be used for change logging. 

60 request.id = uuid.uuid4() 

61 

62 # Apply all registered request processors 

63 with apply_request_processors(request): 

64 response = self.get_response(request) 

65 

66 # Set or renew the language cookie based on the user's preference. This handles two cases: 

67 # 1. The user just logged in (via any auth backend): the user_logged_in signal stores the preferred language on 

68 # the request so we set the cookie here on the login response. 

69 # 2. SESSION_SAVE_EVERY_REQUEST is enabled: renew the language cookie on every request to keep it in sync with 

70 # the session expiry. 

71 if hasattr(request, '_language_cookie'): 71 ↛ 72line 71 didn't jump to line 72 because the condition on line 71 was never true

72 language = request._language_cookie 

73 elif request.user.is_authenticated and settings.SESSION_SAVE_EVERY_REQUEST: 73 ↛ 74line 73 didn't jump to line 74 because the condition on line 73 was never true

74 language = request.user.config.get('locale.language') 

75 else: 

76 language = None 

77 if language: 77 ↛ 78line 77 didn't jump to line 78 because the condition on line 77 was never true

78 response.set_cookie( 

79 key=settings.LANGUAGE_COOKIE_NAME, 

80 value=language, 

81 max_age=request.session.get_expiry_age(), 

82 secure=settings.SESSION_COOKIE_SECURE, 

83 ) 

84 

85 # Attach the unique request ID as an HTTP header. 

86 response['X-Request-ID'] = request.id 

87 

88 # Enable the Vary header to help with caching of HTMX responses 

89 response['Vary'] = 'HX-Request' 

90 

91 # If this is an API request, attach an HTTP header annotating the API version (e.g. '3.5'). 

92 if is_api_request(request): 

93 response['API-Version'] = settings.REST_FRAMEWORK_VERSION 

94 

95 # Clear any cached dynamic config parameters after each request. 

96 clear_config() 

97 

98 return response 

99 

100 def process_exception(self, request, exception): 

101 """ 

102 Implement custom error handling logic for production deployments. 

103 """ 

104 # Don't catch exceptions when in debug mode 

105 if settings.DEBUG: 105 ↛ 106line 105 didn't jump to line 106 because the condition on line 105 was never true

106 return None 

107 

108 # Cleanly handle exceptions that occur from REST or GraphQL API requests 

109 if is_api_request(request) or is_graphql_request(request): 109 ↛ 116line 109 didn't jump to line 116 because the condition on line 109 was always true

110 # Fire Django's got_request_exception signal so error-tracking 

111 # integrations (e.g. Sentry) capture the exception. 

112 got_request_exception.send(sender=self.__class__, request=request) 

113 return handle_rest_api_exception(request) 

114 

115 # Ignore Http404s (defer to Django's built-in 404 handling) 

116 if isinstance(exception, Http404): 

117 return None 

118 

119 # Determine the type of exception. If it's a common issue, return a custom error page with instructions. 

120 custom_template = None 

121 if isinstance(exception, ProgrammingError): 

122 custom_template = 'exceptions/programming_error.html' 

123 elif isinstance(exception, ImportError): 

124 custom_template = 'exceptions/import_error.html' 

125 elif isinstance(exception, PermissionError): 

126 custom_template = 'exceptions/permission_error.html' 

127 

128 # Return a custom error message, or fall back to Django's default 500 error handling 

129 if custom_template: 

130 # Fire Django's got_request_exception signal so error-tracking 

131 # integrations (e.g. Sentry) capture the exception. 

132 got_request_exception.send(sender=self.__class__, request=request) 

133 return handler_500(request, template_name=custom_template) 

134 return None 

135 

136 

137class RemoteUserMiddleware(RemoteUserMiddleware_): 

138 """ 

139 Custom implementation of Django's RemoteUserMiddleware which allows for a user-configurable HTTP header name. 

140 """ 

141 async_capable = False 

142 force_logout_if_no_header = False 

143 

144 def __init__(self, get_response): 

145 if get_response is None: 145 ↛ 146line 145 didn't jump to line 146 because the condition on line 145 was never true

146 raise ValueError("get_response must be provided.") 

147 self.get_response = get_response 

148 

149 @property 

150 def header(self): 

151 return settings.REMOTE_AUTH_HEADER 

152 

153 def __call__(self, request): 

154 logger = logging.getLogger('netbox.authentication.RemoteUserMiddleware') 

155 # Bypass middleware if remote authentication is not enabled 

156 if not settings.REMOTE_AUTH_ENABLED: 156 ↛ 159line 156 didn't jump to line 159 because the condition on line 156 was always true

157 return self.get_response(request) 

158 # AuthenticationMiddleware is required so that request.user exists. 

159 if not hasattr(request, 'user'): 

160 raise ImproperlyConfigured( 

161 "The Django remote user auth middleware requires the" 

162 " authentication middleware to be installed. Edit your" 

163 " MIDDLEWARE setting to insert" 

164 " 'django.contrib.auth.middleware.AuthenticationMiddleware'" 

165 " before the RemoteUserMiddleware class.") 

166 try: 

167 username = request.META[self.header] 

168 except KeyError: 

169 # If specified header doesn't exist then remove any existing 

170 # authenticated remote-user, or return (leaving request.user set to 

171 # AnonymousUser by the AuthenticationMiddleware). 

172 if self.force_logout_if_no_header and request.user.is_authenticated: 

173 self._remove_invalid_user(request) 

174 return self.get_response(request) 

175 # If the user is already authenticated and that user is the user we are 

176 # getting passed in the headers, then the correct user is already 

177 # persisted in the session and we don't need to continue. 

178 if request.user.is_authenticated: 

179 if request.user.get_username() == self.clean_username(username, request): 

180 return self.get_response(request) 

181 # An authenticated user is associated with the request, but 

182 # it does not match the authorized user in the header. 

183 self._remove_invalid_user(request) 

184 

185 # We are seeing this user for the first time in this session, attempt 

186 # to authenticate the user. 

187 if settings.REMOTE_AUTH_GROUP_SYNC_ENABLED: 

188 logger.debug("Trying to sync Groups") 

189 user = auth.authenticate( 

190 request, remote_user=username, remote_groups=self._get_groups(request)) 

191 else: 

192 user = auth.authenticate(request, remote_user=username) 

193 if user: 

194 # User is valid. 

195 # Update the User's Profile if set by request headers 

196 if settings.REMOTE_AUTH_USER_FIRST_NAME in request.META: 

197 user.first_name = request.META[settings.REMOTE_AUTH_USER_FIRST_NAME] 

198 if settings.REMOTE_AUTH_USER_LAST_NAME in request.META: 

199 user.last_name = request.META[settings.REMOTE_AUTH_USER_LAST_NAME] 

200 if settings.REMOTE_AUTH_USER_EMAIL in request.META: 

201 user.email = request.META[settings.REMOTE_AUTH_USER_EMAIL] 

202 user.save() 

203 

204 # Set request.user and persist user in the session 

205 # by logging the user in. 

206 request.user = user 

207 auth.login(request, user) 

208 

209 return self.get_response(request) 

210 

211 def _get_groups(self, request): 

212 logger = logging.getLogger( 

213 'netbox.authentication.RemoteUserMiddleware') 

214 

215 groups_string = request.META.get( 

216 settings.REMOTE_AUTH_GROUP_HEADER, None) 

217 if groups_string: 

218 groups = groups_string.split(settings.REMOTE_AUTH_GROUP_SEPARATOR) 

219 else: 

220 groups = [] 

221 logger.debug(f"Groups are {groups}") 

222 return groups 

223 

224 

225class PrometheusBeforeMiddleware(middleware.PrometheusBeforeMiddleware): 

226 metrics_cls = Metrics 

227 

228 

229class PrometheusAfterMiddleware(middleware.PrometheusAfterMiddleware): 

230 metrics_cls = Metrics 

231 

232 def process_response(self, request, response): 

233 response = super().process_response(request, response) 

234 

235 # Increment REST API request counters 

236 if is_api_request(request): 

237 method = self._method(request) 

238 name = self._get_view_name(request) 

239 self.label_metric(self.metrics.rest_api_requests, request, method=method).inc() 

240 self.label_metric(self.metrics.rest_api_requests_by_view_method, request, method=method, view=name).inc() 

241 

242 # Increment GraphQL API request counters 

243 elif is_graphql_request(request): 

244 self.metrics.graphql_api_requests.inc() 

245 

246 return response 

247 

248 

249class MaintenanceModeMiddleware: 

250 """ 

251 Middleware that checks if the application is in maintenance mode 

252 and restricts write-related operations to the database. 

253 """ 

254 

255 def __init__(self, get_response): 

256 self.get_response = get_response 

257 

258 def __call__(self, request): 

259 if get_config().MAINTENANCE_MODE: 259 ↛ 260line 259 didn't jump to line 260 because the condition on line 259 was never true

260 self._set_session_type( 

261 allow_write=request.path_info.startswith(settings.MAINTENANCE_EXEMPT_PATHS) 

262 ) 

263 

264 return self.get_response(request) 

265 

266 @staticmethod 

267 def _set_session_type(allow_write): 

268 """ 

269 Prevent any write-related database operations. 

270 

271 Args: 

272 allow_write (bool): If True, write operations will be permitted. 

273 """ 

274 with connection.cursor() as cursor: 

275 mode = 'READ WRITE' if allow_write else 'READ ONLY' 

276 cursor.execute(f'SET SESSION CHARACTERISTICS AS TRANSACTION {mode};') 

277 

278 def process_exception(self, request, exception): 

279 """ 

280 Prevent any write-related database operations if an exception is raised. 

281 """ 

282 if get_config().MAINTENANCE_MODE and isinstance(exception, InternalError): 282 ↛ 283line 282 didn't jump to line 283 because the condition on line 282 was never true

283 error_message = 'NetBox is currently operating in maintenance mode and is unable to perform write ' \ 

284 'operations. Please try again later.' 

285 

286 if is_api_request(request) or is_graphql_request(request): 

287 return handle_rest_api_exception(request, error=error_message) 

288 

289 messages.error(request, error_message) 

290 return HttpResponseRedirect(request.path_info) 

291 return None 

292 

293 

294class SocialAuthExceptionMiddleware(SocialAuthExceptionMiddleware_): 

295 """ 

296 Subclass of python-social-auth's exception middleware which surfaces a generic, user-friendly 

297 message rather than exposing the raw social_core exception text to (typically unauthenticated) 

298 users when an SSO/SAML login fails. 

299 """ 

300 def get_message(self, request, exception): 

301 return _("Single sign-on failed. Please try again or contact your administrator.")