summaryrefslogtreecommitdiff
path: root/django/test
diff options
context:
space:
mode:
authorTom Carrick <tom@carrick.eu>2023-11-05 16:41:16 +0100
committerMariusz Felisiak <felisiak.mariusz@gmail.com>2023-11-23 10:39:29 +0100
commita03593967f098cf8dab79065bcabbcebd461f05b (patch)
treef5634ad9ae6d2e42ac8d7254bc85d0d4ceae0bee /django/test
parente76cc93b0168fa3abbafb9af1ab4535814b751f0 (diff)
Fixed #14611 -- Added query_params argument to RequestFactory and Client classes.
Diffstat (limited to 'django/test')
-rw-r--r--django/test/client.py282
1 files changed, 233 insertions, 49 deletions
diff --git a/django/test/client.py b/django/test/client.py
index 766b60863f..48e0588702 100644
--- a/django/test/client.py
+++ b/django/test/client.py
@@ -381,13 +381,22 @@ class RequestFactory:
just as if that view had been hooked up using a URLconf.
"""
- def __init__(self, *, json_encoder=DjangoJSONEncoder, headers=None, **defaults):
+ def __init__(
+ self,
+ *,
+ json_encoder=DjangoJSONEncoder,
+ headers=None,
+ query_params=None,
+ **defaults,
+ ):
self.json_encoder = json_encoder
self.defaults = defaults
self.cookies = SimpleCookie()
self.errors = BytesIO()
if headers:
self.defaults.update(HttpHeaders.to_wsgi_names(headers))
+ if query_params:
+ self.defaults["QUERY_STRING"] = urlencode(query_params, doseq=True)
def _base_environ(self, **request):
"""
@@ -459,18 +468,21 @@ class RequestFactory:
# Refs comment in `get_bytes_from_wsgi()`.
return path.decode("iso-8859-1")
- def get(self, path, data=None, secure=False, *, headers=None, **extra):
+ def get(
+ self, path, data=None, secure=False, *, headers=None, query_params=None, **extra
+ ):
"""Construct a GET request."""
- data = {} if data is None else data
+ if query_params and data:
+ raise ValueError("query_params and data arguments are mutually exclusive.")
+ query_params = data or query_params
+ query_params = {} if query_params is None else query_params
return self.generic(
"GET",
path,
secure=secure,
headers=headers,
- **{
- "QUERY_STRING": urlencode(data, doseq=True),
- **extra,
- },
+ query_params=query_params,
+ **extra,
)
def post(
@@ -481,6 +493,7 @@ class RequestFactory:
secure=False,
*,
headers=None,
+ query_params=None,
**extra,
):
"""Construct a POST request."""
@@ -494,26 +507,37 @@ class RequestFactory:
content_type,
secure=secure,
headers=headers,
+ query_params=query_params,
**extra,
)
- def head(self, path, data=None, secure=False, *, headers=None, **extra):
+ def head(
+ self, path, data=None, secure=False, *, headers=None, query_params=None, **extra
+ ):
"""Construct a HEAD request."""
- data = {} if data is None else data
+ if query_params and data:
+ raise ValueError("query_params and data arguments are mutually exclusive.")
+ query_params = data or query_params
+ query_params = {} if query_params is None else query_params
return self.generic(
"HEAD",
path,
secure=secure,
headers=headers,
- **{
- "QUERY_STRING": urlencode(data, doseq=True),
- **extra,
- },
+ query_params=query_params,
+ **extra,
)
- def trace(self, path, secure=False, *, headers=None, **extra):
+ def trace(self, path, secure=False, *, headers=None, query_params=None, **extra):
"""Construct a TRACE request."""
- return self.generic("TRACE", path, secure=secure, headers=headers, **extra)
+ return self.generic(
+ "TRACE",
+ path,
+ secure=secure,
+ headers=headers,
+ query_params=query_params,
+ **extra,
+ )
def options(
self,
@@ -523,11 +547,19 @@ class RequestFactory:
secure=False,
*,
headers=None,
+ query_params=None,
**extra,
):
"Construct an OPTIONS request."
return self.generic(
- "OPTIONS", path, data, content_type, secure=secure, headers=headers, **extra
+ "OPTIONS",
+ path,
+ data,
+ content_type,
+ secure=secure,
+ headers=headers,
+ query_params=query_params,
+ **extra,
)
def put(
@@ -538,12 +570,20 @@ class RequestFactory:
secure=False,
*,
headers=None,
+ query_params=None,
**extra,
):
"""Construct a PUT request."""
data = self._encode_json(data, content_type)
return self.generic(
- "PUT", path, data, content_type, secure=secure, headers=headers, **extra
+ "PUT",
+ path,
+ data,
+ content_type,
+ secure=secure,
+ headers=headers,
+ query_params=query_params,
+ **extra,
)
def patch(
@@ -554,12 +594,20 @@ class RequestFactory:
secure=False,
*,
headers=None,
+ query_params=None,
**extra,
):
"""Construct a PATCH request."""
data = self._encode_json(data, content_type)
return self.generic(
- "PATCH", path, data, content_type, secure=secure, headers=headers, **extra
+ "PATCH",
+ path,
+ data,
+ content_type,
+ secure=secure,
+ headers=headers,
+ query_params=query_params,
+ **extra,
)
def delete(
@@ -570,12 +618,20 @@ class RequestFactory:
secure=False,
*,
headers=None,
+ query_params=None,
**extra,
):
"""Construct a DELETE request."""
data = self._encode_json(data, content_type)
return self.generic(
- "DELETE", path, data, content_type, secure=secure, headers=headers, **extra
+ "DELETE",
+ path,
+ data,
+ content_type,
+ secure=secure,
+ headers=headers,
+ query_params=query_params,
+ **extra,
)
def generic(
@@ -587,6 +643,7 @@ class RequestFactory:
secure=False,
*,
headers=None,
+ query_params=None,
**extra,
):
"""Construct an arbitrary HTTP request."""
@@ -608,6 +665,8 @@ class RequestFactory:
)
if headers:
extra.update(HttpHeaders.to_wsgi_names(headers))
+ if query_params:
+ extra["QUERY_STRING"] = urlencode(query_params, doseq=True)
r.update(extra)
# If QUERY_STRING is absent or empty, we want to extract it from the URL.
if not r.get("QUERY_STRING"):
@@ -685,6 +744,7 @@ class AsyncRequestFactory(RequestFactory):
secure=False,
*,
headers=None,
+ query_params=None,
**extra,
):
"""Construct an arbitrary HTTP request."""
@@ -705,18 +765,20 @@ class AsyncRequestFactory(RequestFactory):
]
)
s["_body_file"] = FakePayload(data)
- if query_string := extra.pop("QUERY_STRING", None):
+ if query_params:
+ s["query_string"] = urlencode(query_params, doseq=True)
+ elif query_string := extra.pop("QUERY_STRING", None):
s["query_string"] = query_string
+ else:
+ # If QUERY_STRING is absent or empty, we want to extract it from
+ # the URL.
+ s["query_string"] = parsed[4]
if headers:
extra.update(HttpHeaders.to_asgi_names(headers))
s["headers"] += [
(key.lower().encode("ascii"), value.encode("latin1"))
for key, value in extra.items()
]
- # If QUERY_STRING is absent or empty, we want to extract it from the
- # URL.
- if not s.get("query_string"):
- s["query_string"] = parsed[4]
return self.request(**s)
@@ -889,7 +951,14 @@ class ClientMixin:
return response._json
def _follow_redirect(
- self, response, *, data="", content_type="", headers=None, **extra
+ self,
+ response,
+ *,
+ data="",
+ content_type="",
+ headers=None,
+ query_params=None,
+ **extra,
):
"""Follow a single redirect contained in response using GET."""
response_url = response.url
@@ -934,6 +1003,7 @@ class ClientMixin:
content_type=content_type,
follow=False,
headers=headers,
+ query_params=query_params,
**extra,
)
@@ -978,9 +1048,10 @@ class Client(ClientMixin, RequestFactory):
raise_request_exception=True,
*,
headers=None,
+ query_params=None,
**defaults,
):
- super().__init__(headers=headers, **defaults)
+ super().__init__(headers=headers, query_params=query_params, **defaults)
self.handler = ClientHandler(enforce_csrf_checks)
self.raise_request_exception = raise_request_exception
self.exc_info = None
@@ -1042,15 +1113,23 @@ class Client(ClientMixin, RequestFactory):
secure=False,
*,
headers=None,
+ query_params=None,
**extra,
):
"""Request a response from the server using GET."""
self.extra = extra
self.headers = headers
- response = super().get(path, data=data, secure=secure, headers=headers, **extra)
+ response = super().get(
+ path,
+ data=data,
+ secure=secure,
+ headers=headers,
+ query_params=query_params,
+ **extra,
+ )
if follow:
response = self._handle_redirects(
- response, data=data, headers=headers, **extra
+ response, data=data, headers=headers, query_params=query_params, **extra
)
return response
@@ -1063,6 +1142,7 @@ class Client(ClientMixin, RequestFactory):
secure=False,
*,
headers=None,
+ query_params=None,
**extra,
):
"""Request a response from the server using POST."""
@@ -1074,11 +1154,17 @@ class Client(ClientMixin, RequestFactory):
content_type=content_type,
secure=secure,
headers=headers,
+ query_params=query_params,
**extra,
)
if follow:
response = self._handle_redirects(
- response, data=data, content_type=content_type, headers=headers, **extra
+ response,
+ data=data,
+ content_type=content_type,
+ headers=headers,
+ query_params=query_params,
+ **extra,
)
return response
@@ -1090,17 +1176,23 @@ class Client(ClientMixin, RequestFactory):
secure=False,
*,
headers=None,
+ query_params=None,
**extra,
):
"""Request a response from the server using HEAD."""
self.extra = extra
self.headers = headers
response = super().head(
- path, data=data, secure=secure, headers=headers, **extra
+ path,
+ data=data,
+ secure=secure,
+ headers=headers,
+ query_params=query_params,
+ **extra,
)
if follow:
response = self._handle_redirects(
- response, data=data, headers=headers, **extra
+ response, data=data, headers=headers, query_params=query_params, **extra
)
return response
@@ -1113,6 +1205,7 @@ class Client(ClientMixin, RequestFactory):
secure=False,
*,
headers=None,
+ query_params=None,
**extra,
):
"""Request a response from the server using OPTIONS."""
@@ -1124,11 +1217,17 @@ class Client(ClientMixin, RequestFactory):
content_type=content_type,
secure=secure,
headers=headers,
+ query_params=query_params,
**extra,
)
if follow:
response = self._handle_redirects(
- response, data=data, content_type=content_type, headers=headers, **extra
+ response,
+ data=data,
+ content_type=content_type,
+ headers=headers,
+ query_params=query_params,
+ **extra,
)
return response
@@ -1141,6 +1240,7 @@ class Client(ClientMixin, RequestFactory):
secure=False,
*,
headers=None,
+ query_params=None,
**extra,
):
"""Send a resource to the server using PUT."""
@@ -1152,11 +1252,17 @@ class Client(ClientMixin, RequestFactory):
content_type=content_type,
secure=secure,
headers=headers,
+ query_params=query_params,
**extra,
)
if follow:
response = self._handle_redirects(
- response, data=data, content_type=content_type, headers=headers, **extra
+ response,
+ data=data,
+ content_type=content_type,
+ headers=headers,
+ query_params=query_params,
+ **extra,
)
return response
@@ -1169,6 +1275,7 @@ class Client(ClientMixin, RequestFactory):
secure=False,
*,
headers=None,
+ query_params=None,
**extra,
):
"""Send a resource to the server using PATCH."""
@@ -1180,11 +1287,17 @@ class Client(ClientMixin, RequestFactory):
content_type=content_type,
secure=secure,
headers=headers,
+ query_params=query_params,
**extra,
)
if follow:
response = self._handle_redirects(
- response, data=data, content_type=content_type, headers=headers, **extra
+ response,
+ data=data,
+ content_type=content_type,
+ headers=headers,
+ query_params=query_params,
+ **extra,
)
return response
@@ -1197,6 +1310,7 @@ class Client(ClientMixin, RequestFactory):
secure=False,
*,
headers=None,
+ query_params=None,
**extra,
):
"""Send a DELETE request to the server."""
@@ -1208,11 +1322,17 @@ class Client(ClientMixin, RequestFactory):
content_type=content_type,
secure=secure,
headers=headers,
+ query_params=query_params,
**extra,
)
if follow:
response = self._handle_redirects(
- response, data=data, content_type=content_type, headers=headers, **extra
+ response,
+ data=data,
+ content_type=content_type,
+ headers=headers,
+ query_params=query_params,
+ **extra,
)
return response
@@ -1224,17 +1344,23 @@ class Client(ClientMixin, RequestFactory):
secure=False,
*,
headers=None,
+ query_params=None,
**extra,
):
"""Send a TRACE request to the server."""
self.extra = extra
self.headers = headers
response = super().trace(
- path, data=data, secure=secure, headers=headers, **extra
+ path,
+ data=data,
+ secure=secure,
+ headers=headers,
+ query_params=query_params,
+ **extra,
)
if follow:
response = self._handle_redirects(
- response, data=data, headers=headers, **extra
+ response, data=data, headers=headers, query_params=query_params, **extra
)
return response
@@ -1244,6 +1370,7 @@ class Client(ClientMixin, RequestFactory):
data="",
content_type="",
headers=None,
+ query_params=None,
**extra,
):
"""
@@ -1257,6 +1384,7 @@ class Client(ClientMixin, RequestFactory):
data=data,
content_type=content_type,
headers=headers,
+ query_params=query_params,
**extra,
)
response.redirect_chain = redirect_chain
@@ -1278,9 +1406,10 @@ class AsyncClient(ClientMixin, AsyncRequestFactory):
raise_request_exception=True,
*,
headers=None,
+ query_params=None,
**defaults,
):
- super().__init__(headers=headers, **defaults)
+ super().__init__(headers=headers, query_params=query_params, **defaults)
self.handler = AsyncClientHandler(enforce_csrf_checks)
self.raise_request_exception = raise_request_exception
self.exc_info = None
@@ -1341,17 +1470,23 @@ class AsyncClient(ClientMixin, AsyncRequestFactory):
secure=False,
*,
headers=None,
+ query_params=None,
**extra,
):
"""Request a response from the server using GET."""
self.extra = extra
self.headers = headers
response = await super().get(
- path, data=data, secure=secure, headers=headers, **extra
+ path,
+ data=data,
+ secure=secure,
+ headers=headers,
+ query_params=query_params,
+ **extra,
)
if follow:
response = await self._ahandle_redirects(
- response, data=data, headers=headers, **extra
+ response, data=data, headers=headers, query_params=query_params, **extra
)
return response
@@ -1364,6 +1499,7 @@ class AsyncClient(ClientMixin, AsyncRequestFactory):
secure=False,
*,
headers=None,
+ query_params=None,
**extra,
):
"""Request a response from the server using POST."""
@@ -1375,11 +1511,17 @@ class AsyncClient(ClientMixin, AsyncRequestFactory):
content_type=content_type,
secure=secure,
headers=headers,
+ query_params=query_params,
**extra,
)
if follow:
response = await self._ahandle_redirects(
- response, data=data, content_type=content_type, headers=headers, **extra
+ response,
+ data=data,
+ content_type=content_type,
+ headers=headers,
+ query_params=query_params,
+ **extra,
)
return response
@@ -1391,17 +1533,23 @@ class AsyncClient(ClientMixin, AsyncRequestFactory):
secure=False,
*,
headers=None,
+ query_params=None,
**extra,
):
"""Request a response from the server using HEAD."""
self.extra = extra
self.headers = headers
response = await super().head(
- path, data=data, secure=secure, headers=headers, **extra
+ path,
+ data=data,
+ secure=secure,
+ headers=headers,
+ query_params=query_params,
+ **extra,
)
if follow:
response = await self._ahandle_redirects(
- response, data=data, headers=headers, **extra
+ response, data=data, headers=headers, query_params=query_params, **extra
)
return response
@@ -1414,6 +1562,7 @@ class AsyncClient(ClientMixin, AsyncRequestFactory):
secure=False,
*,
headers=None,
+ query_params=None,
**extra,
):
"""Request a response from the server using OPTIONS."""
@@ -1425,11 +1574,17 @@ class AsyncClient(ClientMixin, AsyncRequestFactory):
content_type=content_type,
secure=secure,
headers=headers,
+ query_params=query_params,
**extra,
)
if follow:
response = await self._ahandle_redirects(
- response, data=data, content_type=content_type, headers=headers, **extra
+ response,
+ data=data,
+ content_type=content_type,
+ headers=headers,
+ query_params=query_params,
+ **extra,
)
return response
@@ -1442,6 +1597,7 @@ class AsyncClient(ClientMixin, AsyncRequestFactory):
secure=False,
*,
headers=None,
+ query_params=None,
**extra,
):
"""Send a resource to the server using PUT."""
@@ -1453,11 +1609,17 @@ class AsyncClient(ClientMixin, AsyncRequestFactory):
content_type=content_type,
secure=secure,
headers=headers,
+ query_params=query_params,
**extra,
)
if follow:
response = await self._ahandle_redirects(
- response, data=data, content_type=content_type, headers=headers, **extra
+ response,
+ data=data,
+ content_type=content_type,
+ headers=headers,
+ query_params=query_params,
+ **extra,
)
return response
@@ -1470,6 +1632,7 @@ class AsyncClient(ClientMixin, AsyncRequestFactory):
secure=False,
*,
headers=None,
+ query_params=None,
**extra,
):
"""Send a resource to the server using PATCH."""
@@ -1481,11 +1644,17 @@ class AsyncClient(ClientMixin, AsyncRequestFactory):
content_type=content_type,
secure=secure,
headers=headers,
+ query_params=query_params,
**extra,
)
if follow:
response = await self._ahandle_redirects(
- response, data=data, content_type=content_type, headers=headers, **extra
+ response,
+ data=data,
+ content_type=content_type,
+ headers=headers,
+ query_params=query_params,
+ **extra,
)
return response
@@ -1498,6 +1667,7 @@ class AsyncClient(ClientMixin, AsyncRequestFactory):
secure=False,
*,
headers=None,
+ query_params=None,
**extra,
):
"""Send a DELETE request to the server."""
@@ -1509,11 +1679,17 @@ class AsyncClient(ClientMixin, AsyncRequestFactory):
content_type=content_type,
secure=secure,
headers=headers,
+ query_params=query_params,
**extra,
)
if follow:
response = await self._ahandle_redirects(
- response, data=data, content_type=content_type, headers=headers, **extra
+ response,
+ data=data,
+ content_type=content_type,
+ headers=headers,
+ query_params=query_params,
+ **extra,
)
return response
@@ -1525,17 +1701,23 @@ class AsyncClient(ClientMixin, AsyncRequestFactory):
secure=False,
*,
headers=None,
+ query_params=None,
**extra,
):
"""Send a TRACE request to the server."""
self.extra = extra
self.headers = headers
response = await super().trace(
- path, data=data, secure=secure, headers=headers, **extra
+ path,
+ data=data,
+ secure=secure,
+ headers=headers,
+ query_params=query_params,
+ **extra,
)
if follow:
response = await self._ahandle_redirects(
- response, data=data, headers=headers, **extra
+ response, data=data, headers=headers, query_params=query_params, **extra
)
return response
@@ -1545,6 +1727,7 @@ class AsyncClient(ClientMixin, AsyncRequestFactory):
data="",
content_type="",
headers=None,
+ query_params=None,
**extra,
):
"""
@@ -1558,6 +1741,7 @@ class AsyncClient(ClientMixin, AsyncRequestFactory):
data=data,
content_type=content_type,
headers=headers,
+ query_params=query_params,
**extra,
)
response.redirect_chain = redirect_chain