Loading src/jama/__init__.py +1 −1 Original line number Diff line number Diff line __version__ = "0.6.5" __version__ = "0.6.6" src/jama/middleware.py 0 → 100644 +26 −0 Original line number Diff line number Diff line from django.http import HttpRequest from django.middleware.csrf import CsrfViewMiddleware from resources.models import APIKey class ApiKeyCsrfViewMiddleware(CsrfViewMiddleware): def process_view( self, request: HttpRequest, callback, callback_args, callback_kwargs ): if ( request.method not in ("GET", "HEAD", "OPTIONS", "TRACE") and not getattr(callback, "csrf_exempt", False) and self._request_has_active_api_key(request) ): return self._accept(request) return super().process_view(request, callback, callback_args, callback_kwargs) def _request_has_active_api_key(self, request: HttpRequest) -> bool: api_key = request.headers.get("X-Api-Key") if not api_key: return False return APIKey.objects.filter( key_hash=APIKey.hash_key(api_key), active=True, ).exists() src/jama/settings.py +1 −1 Original line number Diff line number Diff line Loading @@ -485,7 +485,7 @@ MIDDLEWARE = [ "django.contrib.sessions.middleware.SessionMiddleware", "django.middleware.locale.LocaleMiddleware", "django.middleware.common.CommonMiddleware", "django.middleware.csrf.CsrfViewMiddleware", "jama.middleware.ApiKeyCsrfViewMiddleware", "django.contrib.auth.middleware.AuthenticationMiddleware", "django.contrib.messages.middleware.MessageMiddleware", "django.middleware.clickjacking.XFrameOptionsMiddleware", Loading src/jama/tests.py +30 −2 Original line number Diff line number Diff line Loading @@ -3,13 +3,20 @@ from pathlib import Path from django.contrib.auth.models import User from django.db import connection from django.test import RequestFactory, SimpleTestCase, TestCase from django.test import Client, RequestFactory, SimpleTestCase, TestCase from django.test.utils import CaptureQueriesContext from jama import settings as jama_settings from annotations.models import Annotation from jama.iiif import serialize_jama_collection from resources.models import Collection, CollectionMembership, File, FileType, Project from resources.models import ( APIKey, Collection, CollectionMembership, File, FileType, Project, ) class SettingsEnvTemplateTestCase(SimpleTestCase): Loading Loading @@ -87,6 +94,27 @@ class SettingsEnvTemplateTestCase(SimpleTestCase): self.assertIn(f"JAMA_FILES_DIR={var_dir / 'media_source_files'}", content) class ApiKeyCsrfMiddlewareTestCase(TestCase): def setUp(self): self.client = Client(enforce_csrf_checks=True) self.user = User.objects.create(username="api-user") self.api_key = "valid-api-key" api_key, _ = APIKey.create_for_key(self.user, self.api_key) api_key.active = True api_key.save(update_fields=["active"]) def test_valid_api_key_header_skips_csrf_validation(self): response = self.client.post("/status", HTTP_X_API_KEY=self.api_key) self.assertEqual(response.status_code, 200) self.assertEqual(response.content, b"ok") def test_invalid_api_key_header_does_not_skip_csrf_validation(self): response = self.client.post("/status", HTTP_X_API_KEY="invalid") self.assertEqual(response.status_code, 403) class IiifSerializationTestCase(TestCase): def setUp(self): self.factory = RequestFactory() Loading Loading
src/jama/__init__.py +1 −1 Original line number Diff line number Diff line __version__ = "0.6.5" __version__ = "0.6.6"
src/jama/middleware.py 0 → 100644 +26 −0 Original line number Diff line number Diff line from django.http import HttpRequest from django.middleware.csrf import CsrfViewMiddleware from resources.models import APIKey class ApiKeyCsrfViewMiddleware(CsrfViewMiddleware): def process_view( self, request: HttpRequest, callback, callback_args, callback_kwargs ): if ( request.method not in ("GET", "HEAD", "OPTIONS", "TRACE") and not getattr(callback, "csrf_exempt", False) and self._request_has_active_api_key(request) ): return self._accept(request) return super().process_view(request, callback, callback_args, callback_kwargs) def _request_has_active_api_key(self, request: HttpRequest) -> bool: api_key = request.headers.get("X-Api-Key") if not api_key: return False return APIKey.objects.filter( key_hash=APIKey.hash_key(api_key), active=True, ).exists()
src/jama/settings.py +1 −1 Original line number Diff line number Diff line Loading @@ -485,7 +485,7 @@ MIDDLEWARE = [ "django.contrib.sessions.middleware.SessionMiddleware", "django.middleware.locale.LocaleMiddleware", "django.middleware.common.CommonMiddleware", "django.middleware.csrf.CsrfViewMiddleware", "jama.middleware.ApiKeyCsrfViewMiddleware", "django.contrib.auth.middleware.AuthenticationMiddleware", "django.contrib.messages.middleware.MessageMiddleware", "django.middleware.clickjacking.XFrameOptionsMiddleware", Loading
src/jama/tests.py +30 −2 Original line number Diff line number Diff line Loading @@ -3,13 +3,20 @@ from pathlib import Path from django.contrib.auth.models import User from django.db import connection from django.test import RequestFactory, SimpleTestCase, TestCase from django.test import Client, RequestFactory, SimpleTestCase, TestCase from django.test.utils import CaptureQueriesContext from jama import settings as jama_settings from annotations.models import Annotation from jama.iiif import serialize_jama_collection from resources.models import Collection, CollectionMembership, File, FileType, Project from resources.models import ( APIKey, Collection, CollectionMembership, File, FileType, Project, ) class SettingsEnvTemplateTestCase(SimpleTestCase): Loading Loading @@ -87,6 +94,27 @@ class SettingsEnvTemplateTestCase(SimpleTestCase): self.assertIn(f"JAMA_FILES_DIR={var_dir / 'media_source_files'}", content) class ApiKeyCsrfMiddlewareTestCase(TestCase): def setUp(self): self.client = Client(enforce_csrf_checks=True) self.user = User.objects.create(username="api-user") self.api_key = "valid-api-key" api_key, _ = APIKey.create_for_key(self.user, self.api_key) api_key.active = True api_key.save(update_fields=["active"]) def test_valid_api_key_header_skips_csrf_validation(self): response = self.client.post("/status", HTTP_X_API_KEY=self.api_key) self.assertEqual(response.status_code, 200) self.assertEqual(response.content, b"ok") def test_invalid_api_key_header_does_not_skip_csrf_validation(self): response = self.client.post("/status", HTTP_X_API_KEY="invalid") self.assertEqual(response.status_code, 403) class IiifSerializationTestCase(TestCase): def setUp(self): self.factory = RequestFactory() Loading