mirror of
https://github.com/paperless-ngx/paperless-ngx.git
synced 2026-09-02 07:57:15 +00:00
Chore: consolidate pickle hmac signing (#13899)
This commit is contained in:
@@ -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,
|
||||
)
|
||||
|
||||
|
||||
@@ -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:
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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:
|
||||
|
||||
+3
-31
@@ -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,
|
||||
|
||||
@@ -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)
|
||||
@@ -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:
|
||||
|
||||
Reference in New Issue
Block a user