diff --git a/src/documents/caching.py b/src/documents/caching.py index ae697b719..559e83acb 100644 --- a/src/documents/caching.py +++ b/src/documents/caching.py @@ -16,6 +16,9 @@ from django.core.cache import cache from django.core.cache import caches from documents.models import Document +from paperless.signed_pickle import SignedPickleError +from paperless.signed_pickle import signed_pickle_dumps +from paperless.signed_pickle import signed_pickle_loads if TYPE_CHECKING: from django.core.cache.backends.base import BaseCache @@ -118,9 +121,11 @@ class StoredLRUCache(LRUCache): serialized_data = self._backend.get(self._backend_key) try: self._data = ( - pickle.loads(serialized_data) if serialized_data else OrderedDict() + signed_pickle_loads(serialized_data) + if serialized_data + else OrderedDict() ) - except pickle.PickleError: + except (SignedPickleError, pickle.PickleError): logger.warning( "Cache exists in backend but could not be read (possibly invalid format)", ) @@ -132,7 +137,7 @@ class StoredLRUCache(LRUCache): """ self._backend.set( self._backend_key, - pickle.dumps(self._data), + signed_pickle_dumps(self._data), self.backend_ttl, ) diff --git a/src/documents/classifier.py b/src/documents/classifier.py index 519e1eac5..a5c272062 100644 --- a/src/documents/classifier.py +++ b/src/documents/classifier.py @@ -28,6 +28,9 @@ from documents.caching import CLASSIFIER_VERSION_KEY from documents.caching import StoredLRUCache from documents.models import Document from documents.models import MatchingModel +from paperless.signed_pickle import SignedPickleError +from paperless.signed_pickle import signed_pickle_dumps +from paperless.signed_pickle import signed_pickle_loads logger = logging.getLogger("paperless.classifier") @@ -527,10 +530,17 @@ class DocumentClassifier: serialized_result = read_cache.get(key) if serialized_result is None: result = self.data_vectorizer.transform([self.preprocess_content(content)]) - read_cache.set(key, pickle.dumps(result), CACHE_5_MINUTES) + read_cache.set(key, signed_pickle_dumps(result), CACHE_5_MINUTES) else: - read_cache.touch(key, CACHE_5_MINUTES) - result = pickle.loads(serialized_result) + try: + result = signed_pickle_loads(serialized_result) + except SignedPickleError: + result = self.data_vectorizer.transform( + [self.preprocess_content(content)], + ) + read_cache.set(key, signed_pickle_dumps(result), CACHE_5_MINUTES) + else: + read_cache.touch(key, CACHE_5_MINUTES) return result def predict_correspondent(self, content: str) -> int | None: diff --git a/src/documents/tests/test_caching.py b/src/documents/tests/test_caching.py index d75bda3c9..41de11285 100644 --- a/src/documents/tests/test_caching.py +++ b/src/documents/tests/test_caching.py @@ -1,6 +1,7 @@ -import pickle - from documents.caching import StoredLRUCache +from paperless.signed_pickle import HMAC_SIZE +from paperless.signed_pickle import signed_pickle_dumps +from paperless.signed_pickle import signed_pickle_loads def test_lru_cache_entries() -> None: @@ -42,4 +43,16 @@ def test_stored_lru_cache_key_ttl(mocker) -> None: key, data, timeout = mock_backend.set.call_args[0] assert key == "test_key" assert timeout == 321 - assert pickle.loads(data) == {"x": "X", "y": "Y"} + assert signed_pickle_loads(data) == {"x": "X", "y": "Y"} + + +def test_stored_lru_cache_rejects_tampered_data(mocker) -> None: + serialized_data = bytearray(signed_pickle_dumps({"x": "X"})) + serialized_data[HMAC_SIZE] ^= 0xFF + mock_backend = mocker.Mock() + mock_backend.get.return_value = bytes(serialized_data) + cache = StoredLRUCache("test_key", backend=mock_backend) + + cache.load() + + assert cache.get("x") is None diff --git a/src/documents/tests/test_classifier.py b/src/documents/tests/test_classifier.py index 133dc88fe..12dfea4d8 100644 --- a/src/documents/tests/test_classifier.py +++ b/src/documents/tests/test_classifier.py @@ -19,6 +19,8 @@ from documents.models import MatchingModel from documents.models import StoragePath from documents.models import Tag from documents.tests.utils import DirectoriesMixin +from paperless.signed_pickle import HMAC_SIZE +from paperless.signed_pickle import signed_pickle_dumps def dummy_preprocess(content: str, **kwargs): @@ -265,6 +267,27 @@ class TestClassifier(DirectoriesMixin, TestCase): self.assertEqual(mock_preprocess_content.call_count, 2) self.assertEqual(mock_transform.call_count, 2) + def test_vectorize_recomputes_tampered_cache_entry(self) -> None: + cached = bytearray(signed_pickle_dumps(["cached vector"])) + cached[HMAC_SIZE] ^= 0xFF + self.classifier.data_vectorizer = mock.Mock() + self.classifier.data_vectorizer.transform.return_value = ["fresh vector"] + + with ( + mock.patch( + "documents.classifier.read_cache.get", + return_value=bytes(cached), + ), + mock.patch("documents.classifier.read_cache.set") as cache_set, + mock.patch("documents.classifier.read_cache.touch") as cache_touch, + ): + result = self.classifier._vectorize("content") + + self.assertEqual(result, ["fresh vector"]) + self.classifier.data_vectorizer.transform.assert_called_once() + cache_set.assert_called_once() + cache_touch.assert_not_called() + def test_no_retrain_if_no_change(self) -> None: """ GIVEN: diff --git a/src/paperless/celery.py b/src/paperless/celery.py index 3797c840c..e644a71b2 100644 --- a/src/paperless/celery.py +++ b/src/paperless/celery.py @@ -1,12 +1,12 @@ -import hmac import os -import pickle -from hashlib import sha256 from celery import Celery from celery.signals import worker_process_init from kombu.serialization import register +from paperless.signed_pickle import signed_pickle_dumps +from paperless.signed_pickle import signed_pickle_loads + # Set the default Django settings module for the 'celery' program. os.environ.setdefault("DJANGO_SETTINGS_MODULE", "paperless.settings") @@ -18,34 +18,6 @@ os.environ.setdefault("DJANGO_SETTINGS_MODULE", "paperless.settings") # on the worker side using Django's SECRET_KEY. # --------------------------------------------------------------------------- -HMAC_SIZE = 32 # SHA-256 digest length - - -def _get_signing_key() -> bytes: - from django.conf import settings - - return settings.SECRET_KEY.encode() - - -def signed_pickle_dumps(obj: object) -> bytes: - data = pickle.dumps(obj, protocol=pickle.HIGHEST_PROTOCOL) - signature = hmac.new(_get_signing_key(), data, sha256).digest() - return signature + data - - -def signed_pickle_loads(payload: bytes) -> object: - if len(payload) < HMAC_SIZE: - msg = "Signed-pickle payload too short" - raise ValueError(msg) - signature = payload[:HMAC_SIZE] - data = payload[HMAC_SIZE:] - expected = hmac.new(_get_signing_key(), data, sha256).digest() - if not hmac.compare_digest(signature, expected): - msg = "Signed-pickle HMAC verification failed — message may have been tampered with" - raise ValueError(msg) - return pickle.loads(data) - - register( "signed-pickle", signed_pickle_dumps, diff --git a/src/paperless/signed_pickle.py b/src/paperless/signed_pickle.py new file mode 100644 index 000000000..b1f0094aa --- /dev/null +++ b/src/paperless/signed_pickle.py @@ -0,0 +1,39 @@ +from __future__ import annotations + +import hmac +import pickle +from hashlib import sha256 +from typing import Any + +from django.conf import settings + +HMAC_SIZE = sha256().digest_size + + +class SignedPickleError(ValueError): + """Raised when a signed pickle payload cannot be authenticated.""" + + +def _get_signing_key() -> bytes: + return settings.SECRET_KEY.encode() + + +def signed_pickle_dumps(obj: object) -> bytes: + data = pickle.dumps(obj, protocol=pickle.HIGHEST_PROTOCOL) + signature = hmac.new(_get_signing_key(), data, sha256).digest() + return signature + data + + +def signed_pickle_loads(payload: bytes) -> Any: + if len(payload) <= HMAC_SIZE: + msg = "Signed-pickle payload too short" + raise SignedPickleError(msg) + + signature = payload[:HMAC_SIZE] + data = payload[HMAC_SIZE:] + expected = hmac.new(_get_signing_key(), data, sha256).digest() + if not hmac.compare_digest(signature, expected): + msg = "Signed-pickle HMAC verification failed; payload may have been tampered with" + raise SignedPickleError(msg) + + return pickle.loads(data) diff --git a/src/paperless/tests/test_celery.py b/src/paperless/tests/test_celery.py index 364714b6e..83a2e9a84 100644 --- a/src/paperless/tests/test_celery.py +++ b/src/paperless/tests/test_celery.py @@ -6,9 +6,9 @@ from pathlib import Path import pytest from django.test import override_settings -from paperless.celery import HMAC_SIZE from paperless.celery import signed_pickle_dumps from paperless.celery import signed_pickle_loads +from paperless.signed_pickle import HMAC_SIZE class TestSignedPickleSerializer: