diff --git a/api/app/settings/common.py b/api/app/settings/common.py index d948e14b5a9d..4f3690288b85 100644 --- a/api/app/settings/common.py +++ b/api/app/settings/common.py @@ -383,6 +383,9 @@ "whitenoise.middleware.WhiteNoiseMiddleware", "django.contrib.sessions.middleware.SessionMiddleware", "corsheaders.middleware.CorsMiddleware", + # Must come after CorsMiddleware: it can short-circuit with a response of + # its own, and CorsMiddleware needs to wrap it to add CORS headers to that. + "core.middleware.query_params.RejectNulByteQueryParamsMiddleware", "django.middleware.common.CommonMiddleware", "django.middleware.csrf.CsrfViewMiddleware", "django.contrib.auth.middleware.AuthenticationMiddleware", diff --git a/api/core/middleware/query_params.py b/api/core/middleware/query_params.py new file mode 100644 index 000000000000..8eafa65f679f --- /dev/null +++ b/api/core/middleware/query_params.py @@ -0,0 +1,29 @@ +from collections.abc import Callable + +from django.http import HttpRequest, HttpResponse, HttpResponseBadRequest + + +class RejectNulByteQueryParamsMiddleware: + """ + Reject requests whose query parameters contain a NUL (0x00) character. + + Passing one through to a query against the string field of a Postgres + row raises an unhandled `ValueError: A string literal cannot contain + NUL (0x00) characters`, so reject it here, before any view can pass it + to the ORM. + """ + + def __init__(self, get_response: Callable[[HttpRequest], HttpResponse]) -> None: + self.get_response = get_response + + def __call__(self, request: HttpRequest) -> HttpResponse: + # `.values()` only yields the last value per key: a repeated key + # (`?a=x&a=y`) would let a NUL byte in an earlier value slip through. + # `.lists()` yields every value for every key. + if any( + "\x00" in value for _, values in request.GET.lists() for value in values + ): + return HttpResponseBadRequest( + "Query parameters must not contain NUL characters." + ) + return self.get_response(request) diff --git a/api/tests/unit/core/middleware/test_unit_core_middleware_query_params.py b/api/tests/unit/core/middleware/test_unit_core_middleware_query_params.py new file mode 100644 index 000000000000..a663d21289d1 --- /dev/null +++ b/api/tests/unit/core/middleware/test_unit_core_middleware_query_params.py @@ -0,0 +1,63 @@ +from django.http import HttpResponse +from django.test import RequestFactory + +from core.middleware.query_params import RejectNulByteQueryParamsMiddleware + + +def test_reject_nul_byte_query_params_middleware__nul_byte_in_query_param__returns_bad_request( # type: ignore[no-untyped-def] # noqa: E501 + mocker, rf: RequestFactory +): + # Given + mocked_get_response = mocker.MagicMock() + request = rf.get( + "/api/v1/environments/some-key/identities/", {"identifier": "foo\x00bar"} + ) + + middleware = RejectNulByteQueryParamsMiddleware(mocked_get_response) + + # When + response = middleware(request) + + # Then + assert response.status_code == 400 + mocked_get_response.assert_not_called() + + +def test_reject_nul_byte_query_params_middleware__no_nul_byte__calls_get_response( # type: ignore[no-untyped-def] # noqa: E501 + mocker, rf: RequestFactory +): + # Given + a_response = HttpResponse() + mocked_get_response = mocker.MagicMock(return_value=a_response) + request = rf.get( + "/api/v1/environments/some-key/identities/", {"identifier": "foobar"} + ) + + middleware = RejectNulByteQueryParamsMiddleware(mocked_get_response) + + # When + response = middleware(request) + + # Then + assert response is a_response + mocked_get_response.assert_called_once_with(request) + + +def test_reject_nul_byte_query_params_middleware__nul_byte_in_repeated_key__returns_bad_request( # type: ignore[no-untyped-def] # noqa: E501 + mocker, rf: RequestFactory +): + # Given - `identifier` is repeated; `QueryDict.values()` would only see + # the last ("foobar"), silently missing the NUL byte in the first. + mocked_get_response = mocker.MagicMock() + request = rf.get( + "/api/v1/environments/some-key/identities/?identifier=foo%00bar&identifier=foobar" + ) + + middleware = RejectNulByteQueryParamsMiddleware(mocked_get_response) + + # When + response = middleware(request) + + # Then + assert response.status_code == 400 + mocked_get_response.assert_not_called()