Compare commits

...
1 Commits
Author SHA1 Message Date
shamoon 96cffa9ebb Fix: correct text/stream compression workaround 2026-09-10 17:02:16 -07:00
5 changed files with 108 additions and 17 deletions
+36
View File
@@ -38,6 +38,42 @@ class TestChatStreamingViewInputValidation(APITestCase):
) )
assert resp.status_code == status.HTTP_400_BAD_REQUEST assert resp.status_code == status.HTTP_400_BAD_REQUEST
def test_answer_is_not_compressed(self) -> None:
"""
GIVEN:
- A client that accepts compressed responses
WHEN:
- It asks the chat endpoint a question
THEN:
- The answer is streamed unencoded, chunk for chunk
The stream compressors buffer, so a compressed answer arrives in one
piece. The view cannot opt out by flagging the request: DRF's request
wrapper proxies reads but keeps writes to itself, so the flag never
reaches the Django request the middleware sees.
"""
chunks = [f"token{i} " for i in range(40)]
with (
mock.patch(
"documents.views.AIConfig",
return_value=self._mock_ai_enabled(),
),
mock.patch(
"documents.views.stream_chat_with_documents",
return_value=iter(chunks),
),
):
resp = self.client.post(
"/api/documents/chat/",
{"q": "What is in my archive?"},
format="json",
HTTP_ACCEPT_ENCODING="gzip, deflate, br, zstd",
)
assert resp.status_code == status.HTTP_200_OK
assert not resp.has_header("Content-Encoding")
assert list(resp.streaming_content) == [c.encode() for c in chunks]
def test_missing_question_is_rejected(self) -> None: def test_missing_question_is_rejected(self) -> None:
with mock.patch( with mock.patch(
"documents.views.AIConfig", "documents.views.AIConfig",
-1
View File
@@ -2380,7 +2380,6 @@ class ChatStreamingView(GenericAPIView[Any]):
serializer_class = ChatStreamingSerializer serializer_class = ChatStreamingSerializer
def post(self, request, *args, **kwargs): def post(self, request, *args, **kwargs):
request.compress_exempt = True
ai_config = AIConfig() ai_config = AIConfig()
if not ai_config.ai_enabled: if not ai_config.ai_enabled:
return HttpResponseBadRequest("AI is required for this feature") return HttpResponseBadRequest("AI is required for this feature")
+15
View File
@@ -1,8 +1,23 @@
from compression_middleware.middleware import CompressionMiddleware
from django.conf import settings from django.conf import settings
from paperless import version from paperless import version
class StreamAwareCompressionMiddleware(CompressionMiddleware):
"""
Bypasses compression for server-sent streams (text/event-stream).
See https://github.com/friedelwolff/django-compression-middleware/pull/7
"""
def process_response(self, request, response):
content_type = response.headers.get("Content-Type", "")
if content_type.startswith("text/event-stream"):
return response
return super().process_response(request, response)
class ApiVersionMiddleware: class ApiVersionMiddleware:
def __init__(self, get_response): def __init__(self, get_response):
self.get_response = get_response self.get_response = get_response
+3 -16
View File
@@ -10,7 +10,6 @@ from pathlib import Path
from typing import Final from typing import Final
from urllib.parse import urlparse from urllib.parse import urlparse
from compression_middleware.middleware import CompressionMiddleware
from django.core.exceptions import ImproperlyConfigured from django.core.exceptions import ImproperlyConfigured
from django.utils.translation import gettext_lazy as _ from django.utils.translation import gettext_lazy as _
from dotenv import load_dotenv from dotenv import load_dotenv
@@ -201,22 +200,10 @@ MIDDLEWARE = [
"allauth.account.middleware.AccountMiddleware", "allauth.account.middleware.AccountMiddleware",
] ]
# Optional to enable compression # Optional to enable compression. The subclass leaves server-sent events
# uncompressed; see paperless.middleware.StreamAwareCompressionMiddleware.
if get_bool_from_env("PAPERLESS_ENABLE_COMPRESSION", "yes"): # pragma: no cover if get_bool_from_env("PAPERLESS_ENABLE_COMPRESSION", "yes"): # pragma: no cover
MIDDLEWARE.insert(0, "compression_middleware.middleware.CompressionMiddleware") MIDDLEWARE.insert(0, "paperless.middleware.StreamAwareCompressionMiddleware")
# Workaround to not compress streaming responses (e.g. chat).
# See https://github.com/friedelwolff/django-compression-middleware/pull/7
original_process_response = CompressionMiddleware.process_response
def patched_process_response(self, request, response):
if getattr(request, "compress_exempt", False):
return response
return original_process_response(self, request, response)
CompressionMiddleware.process_response = patched_process_response
ROOT_URLCONF = "paperless.urls" ROOT_URLCONF = "paperless.urls"
@@ -0,0 +1,54 @@
from django.http import HttpResponse
from django.http import StreamingHttpResponse
from django.test import RequestFactory
from django.test import TestCase
from paperless.middleware import StreamAwareCompressionMiddleware
class TestStreamAwareCompressionMiddleware(TestCase):
def setUp(self) -> None:
super().setUp()
self.factory = RequestFactory()
self.middleware = StreamAwareCompressionMiddleware(lambda request: None)
def _request(self):
return self.factory.get(
"/api/documents/chat/",
HTTP_ACCEPT_ENCODING="gzip, deflate, br, zstd",
)
def test_event_stream_is_not_compressed(self) -> None:
"""
GIVEN:
- A server-sent event response produced chunk by chunk
WHEN:
- The compression middleware processes it
THEN:
- It is passed through unencoded, one wire chunk per source chunk
"""
chunks = [f"token{i} ".encode() for i in range(40)]
response = StreamingHttpResponse(
iter(chunks),
content_type="text/event-stream",
)
response = self.middleware.process_response(self._request(), response)
assert not response.has_header("Content-Encoding")
assert list(response.streaming_content) == chunks
def test_regular_response_is_still_compressed(self) -> None:
"""
GIVEN:
- An ordinary response large enough to be worth compressing
WHEN:
- The compression middleware processes it
THEN:
- It is compressed as before
"""
response = HttpResponse(b"a" * 5000, content_type="application/json")
response = self.middleware.process_response(self._request(), response)
assert response.has_header("Content-Encoding")