testclient.py 28 KB

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