mirror of
https://github.com/paperless-ngx/paperless-ngx.git
synced 2026-09-11 12:18:02 +00:00
Compare commits
1
Commits
dev
...
fix/drf-stream
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
96cffa9ebb |
@@ -38,6 +38,42 @@ class TestChatStreamingViewInputValidation(APITestCase):
|
||||
)
|
||||
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:
|
||||
with mock.patch(
|
||||
"documents.views.AIConfig",
|
||||
|
||||
@@ -2380,7 +2380,6 @@ class ChatStreamingView(GenericAPIView[Any]):
|
||||
serializer_class = ChatStreamingSerializer
|
||||
|
||||
def post(self, request, *args, **kwargs):
|
||||
request.compress_exempt = True
|
||||
ai_config = AIConfig()
|
||||
if not ai_config.ai_enabled:
|
||||
return HttpResponseBadRequest("AI is required for this feature")
|
||||
|
||||
@@ -1,8 +1,23 @@
|
||||
from compression_middleware.middleware import CompressionMiddleware
|
||||
from django.conf import settings
|
||||
|
||||
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:
|
||||
def __init__(self, get_response):
|
||||
self.get_response = get_response
|
||||
|
||||
@@ -10,7 +10,6 @@ from pathlib import Path
|
||||
from typing import Final
|
||||
from urllib.parse import urlparse
|
||||
|
||||
from compression_middleware.middleware import CompressionMiddleware
|
||||
from django.core.exceptions import ImproperlyConfigured
|
||||
from django.utils.translation import gettext_lazy as _
|
||||
from dotenv import load_dotenv
|
||||
@@ -201,22 +200,10 @@ MIDDLEWARE = [
|
||||
"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
|
||||
MIDDLEWARE.insert(0, "compression_middleware.middleware.CompressionMiddleware")
|
||||
|
||||
# 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
|
||||
MIDDLEWARE.insert(0, "paperless.middleware.StreamAwareCompressionMiddleware")
|
||||
|
||||
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")
|
||||
Reference in New Issue
Block a user