Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
3 changes: 3 additions & 0 deletions api/app/settings/common.py
Original file line number Diff line number Diff line change
Expand Up @@ -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",
Expand Down
29 changes: 29 additions & 0 deletions api/core/middleware/query_params.py
Original file line number Diff line number Diff line change
@@ -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)
Comment thread
coderabbitai[bot] marked this conversation as resolved.
Original file line number Diff line number Diff line change
@@ -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()
Loading