diff --git a/app/__init__.py b/app/__init__.py index 7169dcd27..c3d48752c 100644 --- a/app/__init__.py +++ b/app/__init__.py @@ -19,6 +19,7 @@ from notifications_utils.clients.encryption.encryption_client import Encryption from notifications_utils import logging, request_helper from sqlalchemy import event from werkzeug.exceptions import HTTPException as WerkzeugHTTPException +from werkzeug.exceptions import RequestEntityTooLarge as WerkzeugRequestEntityTooLarge from werkzeug.local import LocalProxy from app.celery.celery import NotifyCelery @@ -281,6 +282,15 @@ def init_app(app): g.start = monotonic() g.endpoint = request.endpoint + @app.before_request + def check_content_length(): + if ( + request.content_length is not None + and current_app.config['MAX_CONTENT_LENGTH'] is not None + and request.content_length > current_app.config['MAX_CONTENT_LENGTH'] + ): + raise WerkzeugRequestEntityTooLarge() + @app.after_request def after_request(response): CONCURRENT_REQUESTS.dec() diff --git a/tests/app/test_request_size.py b/tests/app/test_request_size.py new file mode 100644 index 000000000..650a181c6 --- /dev/null +++ b/tests/app/test_request_size.py @@ -0,0 +1,36 @@ +import pytest +import json + +@pytest.mark.parametrize('endpoint, max_content_length, expected_status_code', [ + ("/_status", 5*1024*1024, 413), + ("/provider-details", 5*1024*1024, 413), + ("/v2/notifications/email", 5*1024*1024, 413), + + ("/_status", None, 200), + ("/provider-details", None, 405), + ("/v2/notifications/email", None, 401), +]) +def test_request_status_when_content_length_is_set( + notify_api, + sample_email_template_with_placeholders, + mocker, + endpoint, + max_content_length, + expected_status_code): + + notify_api.config['MAX_CONTENT_LENGTH'] = max_content_length + large_name = "J" * (max_content_length or 1 + 1) + data = { + 'email_address': 'ok@ok.com', + 'template_id': str(sample_email_template_with_placeholders.id), + 'personalisation': { + 'name': large_name + } + } + with notify_api.test_client() as client: + response = client.post( + path=endpoint, + data=json.dumps(data), + headers=[('Content-Type', 'application/json')]) + + assert response.status_code == expected_status_code