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
|
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",
|
||||||
|
|||||||
@@ -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")
|
||||||
|
|||||||
@@ -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
|
||||||
|
|||||||
@@ -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")
|
||||||
Reference in New Issue
Block a user