Coverage for /pythoncovmergedfiles/medio/medio/usr/local/lib/python3.11/site-packages/starlette/testclient.py: 23%
Shortcuts on this page
r m x toggle line displays
j k next/prev highlighted chunk
0 (zero) top of page
1 (one) first highlighted chunk
Shortcuts on this page
r m x toggle line displays
j k next/prev highlighted chunk
0 (zero) top of page
1 (one) first highlighted chunk
1from __future__ import annotations
3import contextlib
4import inspect
5import io
6import json
7import math
8import sys
9import warnings
10from collections.abc import Awaitable, Callable, Generator, Iterable, Mapping, Sequence
11from concurrent.futures import Future
12from contextlib import AbstractContextManager
13from types import GeneratorType
14from typing import TYPE_CHECKING, Any, Literal, TypedDict, TypeGuard, cast
15from urllib.parse import unquote, urljoin
17import anyio
18import anyio.abc
19import anyio.from_thread
20from anyio.streams.stapled import StapledObjectStream
22from starlette._utils import is_async_callable
23from starlette.exceptions import StarletteDeprecationWarning
24from starlette.types import ASGIApp, Message, Receive, Scope, Send
25from starlette.websockets import WebSocketDisconnect
27if sys.version_info >= (3, 11): # pragma: no cover
28 from typing import Self
29else: # pragma: no cover
30 from typing_extensions import Self
32if TYPE_CHECKING:
33 import httpx2 as httpx
34else:
35 try:
36 import httpx2 as httpx
37 except ModuleNotFoundError: # pragma: no cover
38 try:
39 import httpx
40 except ModuleNotFoundError:
41 raise RuntimeError(
42 "The starlette.testclient module requires the httpx2 package to be installed.\n"
43 "You can install this with:\n"
44 " $ pip install httpx2\n"
45 ) from None
46 else:
47 warnings.warn(
48 "Using `httpx` with `starlette.testclient` is deprecated; install `httpx2` instead.",
49 StarletteDeprecationWarning,
50 stacklevel=2,
51 )
53_PortalFactoryType = Callable[[], AbstractContextManager[anyio.from_thread.BlockingPortal]]
55ASGIInstance = Callable[[Receive, Send], Awaitable[None]]
56ASGI2App = Callable[[Scope], ASGIInstance]
57ASGI3App = Callable[[Scope, Receive, Send], Awaitable[None]]
60_RequestData = Mapping[str, str | Iterable[str] | bytes]
63def _is_asgi3(app: ASGI2App | ASGI3App) -> TypeGuard[ASGI3App]:
64 if inspect.isclass(app):
65 return hasattr(app, "__await__")
66 return is_async_callable(app)
69class _WrapASGI2:
70 """
71 Provide an ASGI3 interface onto an ASGI2 app.
72 """
74 def __init__(self, app: ASGI2App) -> None:
75 self.app = app
77 async def __call__(self, scope: Scope, receive: Receive, send: Send) -> None:
78 instance = self.app(scope)
79 await instance(receive, send)
82class _AsyncBackend(TypedDict):
83 backend: str
84 backend_options: dict[str, Any]
87class _Upgrade(Exception):
88 def __init__(self, session: WebSocketTestSession) -> None:
89 self.session = session
92class WebSocketDenialResponse( # type: ignore[misc]
93 httpx.Response,
94 WebSocketDisconnect,
95):
96 """
97 A special case of `WebSocketDisconnect`, raised in the `TestClient` if the
98 `WebSocket` is closed before being accepted with a `send_denial_response()`.
99 """
102class WebSocketTestSession:
103 def __init__(
104 self,
105 app: ASGI3App,
106 scope: Scope,
107 portal_factory: _PortalFactoryType,
108 ) -> None:
109 self.app = app
110 self.scope = scope
111 self.accepted_subprotocol = None
112 self.portal_factory = portal_factory
113 self.extra_headers = None
115 def __enter__(self) -> Self:
116 with contextlib.ExitStack() as stack:
117 self.portal = portal = stack.enter_context(self.portal_factory())
118 fut, cs = portal.start_task(self._run)
119 stack.callback(fut.result)
120 stack.callback(portal.call, cs.cancel)
121 self.send({"type": "websocket.connect"})
122 message = self.receive()
123 self._raise_on_close(message)
124 self.accepted_subprotocol = message.get("subprotocol", None)
125 self.extra_headers = message.get("headers", None)
126 stack.callback(self.close, 1000)
127 self.exit_stack = stack.pop_all()
128 return self
130 def __exit__(self, *args: Any) -> bool | None:
131 return self.exit_stack.__exit__(*args)
133 async def _run(self, *, task_status: anyio.abc.TaskStatus[anyio.CancelScope]) -> None:
134 """
135 The sub-thread in which the websocket session runs.
136 """
137 send_tx, send_rx = anyio.create_memory_object_stream[Message](math.inf)
138 receive_tx, receive_rx = anyio.create_memory_object_stream[Message](math.inf)
139 with send_tx, send_rx, receive_tx, receive_rx, anyio.CancelScope() as cs:
140 self._receive_tx = receive_tx
141 self._send_rx = send_rx
142 task_status.started(cs)
143 await self.app(self.scope, receive_rx.receive, send_tx.send)
145 # wait for cs.cancel to be called before closing streams
146 await anyio.sleep_forever()
148 def _raise_on_close(self, message: Message) -> None:
149 if message["type"] == "websocket.close":
150 raise WebSocketDisconnect(code=message.get("code", 1000), reason=message.get("reason", ""))
151 elif message["type"] == "websocket.http.response.start":
152 status_code: int = message["status"]
153 headers: list[tuple[bytes, bytes]] = message["headers"]
154 body: list[bytes] = []
155 while True:
156 message = self.receive()
157 assert message["type"] == "websocket.http.response.body"
158 body.append(message["body"])
159 if not message.get("more_body", False):
160 break
161 raise WebSocketDenialResponse(status_code=status_code, headers=headers, content=b"".join(body))
163 def send(self, message: Message) -> None:
164 self.portal.call(self._receive_tx.send, message)
166 def send_text(self, data: str) -> None:
167 self.send({"type": "websocket.receive", "text": data})
169 def send_bytes(self, data: bytes) -> None:
170 self.send({"type": "websocket.receive", "bytes": data})
172 def send_json(self, data: Any, mode: Literal["text", "binary"] = "text") -> None:
173 text = json.dumps(data, separators=(",", ":"), ensure_ascii=False)
174 if mode == "text":
175 self.send({"type": "websocket.receive", "text": text})
176 else:
177 self.send({"type": "websocket.receive", "bytes": text.encode("utf-8")})
179 def close(self, code: int = 1000, reason: str | None = None) -> None:
180 self.send({"type": "websocket.disconnect", "code": code, "reason": reason})
182 def receive(self) -> Message:
183 return self.portal.call(self._send_rx.receive)
185 def receive_text(self) -> str:
186 message = self.receive()
187 self._raise_on_close(message)
188 return cast(str, message["text"])
190 def receive_bytes(self) -> bytes:
191 message = self.receive()
192 self._raise_on_close(message)
193 return cast(bytes, message["bytes"])
195 def receive_json(self, mode: Literal["text", "binary"] = "text") -> Any:
196 message = self.receive()
197 self._raise_on_close(message)
198 if mode == "text":
199 text = message["text"]
200 else:
201 text = message["bytes"].decode("utf-8")
202 return json.loads(text)
205class _TestClientTransport(httpx.BaseTransport):
206 def __init__(
207 self,
208 app: ASGI3App,
209 portal_factory: _PortalFactoryType,
210 raise_server_exceptions: bool = True,
211 root_path: str = "",
212 *,
213 client: tuple[str, int],
214 app_state: dict[str, Any],
215 ) -> None:
216 self.app = app
217 self.raise_server_exceptions = raise_server_exceptions
218 self.root_path = root_path
219 self.portal_factory = portal_factory
220 self.app_state = app_state
221 self.client = client
223 def handle_request(self, request: httpx.Request) -> httpx.Response:
224 scheme = request.url.scheme
225 host = request.url.raw_host.decode(encoding="ascii")
226 path = request.url.path
227 raw_path = request.url.raw_path
228 query = request.url.query.decode(encoding="ascii")
230 default_port = {"http": 80, "ws": 80, "https": 443, "wss": 443}[scheme]
231 port = request.url.port
232 if port is None:
233 port = default_port
235 # Include the 'host' header.
236 if "host" in request.headers:
237 headers: list[tuple[bytes, bytes]] = []
238 else: # pragma: no cover
239 headers = [(b"host", request.url.netloc)]
241 # Include other request headers.
242 headers += [(key.lower().encode(), value.encode()) for key, value in request.headers.multi_items()]
244 scope: dict[str, Any]
246 if scheme in {"ws", "wss"}:
247 subprotocol = request.headers.get("sec-websocket-protocol", None)
248 if subprotocol is None:
249 subprotocols: Sequence[str] = []
250 else:
251 subprotocols = [value.strip() for value in subprotocol.split(",")]
252 scope = {
253 "type": "websocket",
254 "path": unquote(path),
255 "raw_path": raw_path.split(b"?", 1)[0],
256 "root_path": self.root_path,
257 "scheme": scheme,
258 "query_string": query.encode(),
259 "headers": headers,
260 "client": self.client,
261 "server": [host, port],
262 "subprotocols": subprotocols,
263 "state": self.app_state.copy(),
264 "extensions": {"websocket.http.response": {}},
265 }
266 session = WebSocketTestSession(self.app, scope, self.portal_factory)
267 raise _Upgrade(session)
269 scope = {
270 "type": "http",
271 "http_version": "1.1",
272 "method": request.method,
273 "path": unquote(path),
274 "raw_path": raw_path.split(b"?", 1)[0],
275 "root_path": self.root_path,
276 "scheme": scheme,
277 "query_string": query.encode(),
278 "headers": headers,
279 "client": self.client,
280 "server": [host, port],
281 "extensions": {"http.response.debug": {}, "http.response.trailers": {}},
282 "state": self.app_state.copy(),
283 }
285 request_complete = False
286 response_started = False
287 body_complete = False
288 trailers_expected = False
289 trailers: list[tuple[bytes, bytes]] = []
290 response_complete: anyio.Event
291 raw_kwargs: dict[str, Any] = {"stream": io.BytesIO()}
292 debug_info: dict[str, Any] | None = None
294 async def receive() -> Message:
295 nonlocal request_complete
297 if request_complete:
298 if not response_complete.is_set():
299 await response_complete.wait()
300 return {"type": "http.disconnect"}
302 body = request.read()
303 if isinstance(body, str):
304 body_bytes: bytes = body.encode("utf-8") # pragma: no cover
305 elif body is None:
306 body_bytes = b"" # pragma: no cover
307 elif isinstance(body, GeneratorType):
308 try: # pragma: no cover
309 chunk = body.send(None)
310 if isinstance(chunk, str):
311 chunk = chunk.encode("utf-8")
312 return {"type": "http.request", "body": chunk, "more_body": True}
313 except StopIteration: # pragma: no cover
314 request_complete = True
315 return {"type": "http.request", "body": b""}
316 else:
317 body_bytes = body
319 request_complete = True
320 return {"type": "http.request", "body": body_bytes}
322 async def send(message: Message) -> None:
323 nonlocal raw_kwargs, response_started, debug_info, body_complete, trailers_expected
325 if message["type"] == "http.response.start":
326 assert not response_started, 'Received multiple "http.response.start" messages.'
327 raw_kwargs["status_code"] = message["status"]
328 raw_kwargs["headers"] = [(key.decode(), value.decode()) for key, value in message.get("headers", [])]
329 response_started = True
330 trailers_expected = message.get("trailers", False)
331 elif message["type"] == "http.response.body":
332 assert response_started, 'Received "http.response.body" without "http.response.start".'
333 assert not response_complete.is_set(), 'Received "http.response.body" after response completed.'
334 assert not body_complete, 'Received "http.response.body" after body completed.'
335 body = message.get("body", b"")
336 more_body = message.get("more_body", False)
337 if request.method != "HEAD":
338 raw_kwargs["stream"].write(body)
339 if not more_body:
340 raw_kwargs["stream"].seek(0)
341 body_complete = True
342 if not trailers_expected:
343 response_complete.set()
344 elif message["type"] == "http.response.trailers":
345 assert trailers_expected, 'Received "http.response.trailers" without declaring trailers.'
346 assert body_complete, 'Received "http.response.trailers" before body completed.'
347 assert not response_complete.is_set(), 'Received "http.response.trailers" after response completed.'
348 trailers.extend(message.get("headers", []))
349 if not message.get("more_trailers", False):
350 response_complete.set()
351 elif message["type"] == "http.response.debug":
352 debug_info = message["info"]
354 try:
355 with self.portal_factory() as portal:
356 response_complete = portal.call(anyio.Event)
357 portal.call(self.app, scope, receive, send)
358 except BaseException as exc:
359 if self.raise_server_exceptions:
360 raise exc
362 if self.raise_server_exceptions:
363 assert response_started, "TestClient did not receive any response."
364 elif not response_started:
365 raw_kwargs = {
366 "status_code": 500,
367 "headers": [],
368 "stream": io.BytesIO(),
369 }
371 raw_kwargs["stream"] = httpx.ByteStream(raw_kwargs["stream"].read())
373 response = httpx.Response(**raw_kwargs, request=request)
374 if trailers_expected:
375 response.extensions["http.response.trailers"] = trailers
376 if debug_info is not None:
377 response.extensions["http.response.debug"] = debug_info
378 if "template" in debug_info:
379 response.template = debug_info["template"] # type: ignore[attr-defined]
380 if "context" in debug_info:
381 response.context = debug_info["context"] # type: ignore[attr-defined]
382 return response
385class TestClient(httpx.Client):
386 __test__ = False
387 task: Future[None]
388 portal: anyio.from_thread.BlockingPortal | None = None
390 def __init__(
391 self,
392 app: ASGIApp,
393 base_url: str = "http://testserver",
394 raise_server_exceptions: bool = True,
395 root_path: str = "",
396 backend: Literal["asyncio", "trio"] = "asyncio",
397 backend_options: dict[str, Any] | None = None,
398 cookies: httpx._types.CookieTypes | None = None,
399 headers: dict[str, str] | None = None,
400 follow_redirects: bool = True,
401 client: tuple[str, int] = ("testclient", 50000),
402 ) -> None:
403 self.async_backend = _AsyncBackend(backend=backend, backend_options=backend_options or {})
404 if _is_asgi3(app):
405 asgi_app = app
406 else:
407 app = cast(ASGI2App, app) # type: ignore[assignment]
408 asgi_app = _WrapASGI2(app) # type: ignore[arg-type]
409 self.app = asgi_app
410 self.app_state: dict[str, Any] = {}
411 transport = _TestClientTransport(
412 self.app,
413 portal_factory=self._portal_factory,
414 raise_server_exceptions=raise_server_exceptions,
415 root_path=root_path,
416 app_state=self.app_state,
417 client=client,
418 )
419 if headers is None:
420 headers = {}
421 headers.setdefault("user-agent", "testclient")
422 super().__init__(
423 base_url=base_url,
424 headers=headers,
425 transport=transport,
426 follow_redirects=follow_redirects,
427 cookies=cookies,
428 )
430 @contextlib.contextmanager
431 def _portal_factory(self) -> Generator[anyio.from_thread.BlockingPortal, None, None]:
432 if self.portal is not None:
433 yield self.portal
434 else:
435 with anyio.from_thread.start_blocking_portal(**self.async_backend) as portal:
436 yield portal
438 def request( # type: ignore[override]
439 self,
440 method: str,
441 url: httpx._types.URLTypes,
442 *,
443 content: httpx._types.RequestContent | None = None,
444 data: _RequestData | None = None,
445 files: httpx._types.RequestFiles | None = None,
446 json: Any = None,
447 params: httpx._types.QueryParamTypes | None = None,
448 headers: httpx._types.HeaderTypes | None = None,
449 cookies: httpx._types.CookieTypes | None = None,
450 auth: httpx._types.AuthTypes | httpx._client.UseClientDefault = httpx._client.USE_CLIENT_DEFAULT,
451 follow_redirects: bool | httpx._client.UseClientDefault = httpx._client.USE_CLIENT_DEFAULT,
452 timeout: httpx._types.TimeoutTypes | httpx._client.UseClientDefault = httpx._client.USE_CLIENT_DEFAULT,
453 extensions: dict[str, Any] | None = None,
454 ) -> httpx.Response:
455 if timeout is not httpx.USE_CLIENT_DEFAULT:
456 warnings.warn(
457 "You should not use the 'timeout' argument with the TestClient. "
458 "See https://github.com/Kludex/starlette/issues/1108 for more information.",
459 StarletteDeprecationWarning,
460 stacklevel=2,
461 )
462 url = self._merge_url(url)
463 return super().request(
464 method,
465 url,
466 content=content,
467 data=data,
468 files=files,
469 json=json,
470 params=params,
471 headers=headers,
472 cookies=cookies,
473 auth=auth,
474 follow_redirects=follow_redirects,
475 timeout=timeout,
476 extensions=extensions,
477 )
479 def websocket_connect(
480 self,
481 url: str,
482 subprotocols: Sequence[str] | None = None,
483 **kwargs: Any,
484 ) -> WebSocketTestSession:
485 url = urljoin("ws://testserver", url)
486 headers = kwargs.get("headers", {})
487 headers.setdefault("connection", "upgrade")
488 headers.setdefault("sec-websocket-key", "testserver==")
489 headers.setdefault("sec-websocket-version", "13")
490 if subprotocols is not None:
491 headers.setdefault("sec-websocket-protocol", ", ".join(subprotocols))
492 kwargs["headers"] = headers
493 try:
494 super().request("GET", url, **kwargs)
495 except _Upgrade as exc:
496 session = exc.session
497 else:
498 raise RuntimeError("Expected WebSocket upgrade") # pragma: no cover
500 return session
502 def __enter__(self) -> Self:
503 with contextlib.ExitStack() as stack:
504 self.portal = portal = stack.enter_context(anyio.from_thread.start_blocking_portal(**self.async_backend))
506 @stack.callback
507 def reset_portal() -> None:
508 self.portal = None
510 send = anyio.create_memory_object_stream[Message | None](math.inf)
511 receive = anyio.create_memory_object_stream[Message](math.inf)
512 for channel in (*send, *receive):
513 stack.callback(channel.close)
514 self.stream_send = StapledObjectStream(*send)
515 self.stream_receive = StapledObjectStream(*receive)
516 self.task = portal.start_task_soon(self.lifespan)
517 portal.call(self.wait_startup)
519 @stack.callback
520 def wait_shutdown() -> None:
521 portal.call(self.wait_shutdown)
523 self.exit_stack = stack.pop_all()
525 return self
527 def __exit__(self, *args: Any) -> None:
528 self.exit_stack.close()
530 async def lifespan(self) -> None:
531 scope = {"type": "lifespan", "state": self.app_state}
532 try:
533 await self.app(scope, self.stream_receive.receive, self.stream_send.send)
534 finally:
535 await self.stream_send.send(None)
537 async def wait_startup(self) -> None:
538 await self.stream_receive.send({"type": "lifespan.startup"})
540 async def receive() -> Any:
541 message = await self.stream_send.receive()
542 if message is None:
543 self.task.result()
544 return message
546 message = await receive()
547 assert message["type"] in (
548 "lifespan.startup.complete",
549 "lifespan.startup.failed",
550 )
551 if message["type"] == "lifespan.startup.failed":
552 await receive()
554 async def wait_shutdown(self) -> None:
555 async def receive() -> Any:
556 message = await self.stream_send.receive()
557 if message is None:
558 self.task.result()
559 return message
561 await self.stream_receive.send({"type": "lifespan.shutdown"})
562 message = await receive()
563 assert message["type"] in (
564 "lifespan.shutdown.complete",
565 "lifespan.shutdown.failed",
566 )
567 if message["type"] == "lifespan.shutdown.failed":
568 await receive()