diff options
| author | Tom Carrick <tom@carrick.eu> | 2023-11-05 16:41:16 +0100 |
|---|---|---|
| committer | Mariusz Felisiak <felisiak.mariusz@gmail.com> | 2023-11-23 10:39:29 +0100 |
| commit | a03593967f098cf8dab79065bcabbcebd461f05b (patch) | |
| tree | f5634ad9ae6d2e42ac8d7254bc85d0d4ceae0bee /django/test | |
| parent | e76cc93b0168fa3abbafb9af1ab4535814b751f0 (diff) | |
Fixed #14611 -- Added query_params argument to RequestFactory and Client classes.
Diffstat (limited to 'django/test')
| -rw-r--r-- | django/test/client.py | 282 |
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 |
