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

314 statements  

1from __future__ import annotations 

2 

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 

16 

17import anyio 

18import anyio.abc 

19import anyio.from_thread 

20from anyio.streams.stapled import StapledObjectStream 

21 

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 

26 

27if sys.version_info >= (3, 11): # pragma: no cover 

28 from typing import Self 

29else: # pragma: no cover 

30 from typing_extensions import Self 

31 

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 ) 

52 

53_PortalFactoryType = Callable[[], AbstractContextManager[anyio.from_thread.BlockingPortal]] 

54 

55ASGIInstance = Callable[[Receive, Send], Awaitable[None]] 

56ASGI2App = Callable[[Scope], ASGIInstance] 

57ASGI3App = Callable[[Scope, Receive, Send], Awaitable[None]] 

58 

59 

60_RequestData = Mapping[str, str | Iterable[str] | bytes] 

61 

62 

63def _is_asgi3(app: ASGI2App | ASGI3App) -> TypeGuard[ASGI3App]: 

64 if inspect.isclass(app): 

65 return hasattr(app, "__await__") 

66 return is_async_callable(app) 

67 

68 

69class _WrapASGI2: 

70 """ 

71 Provide an ASGI3 interface onto an ASGI2 app. 

72 """ 

73 

74 def __init__(self, app: ASGI2App) -> None: 

75 self.app = app 

76 

77 async def __call__(self, scope: Scope, receive: Receive, send: Send) -> None: 

78 instance = self.app(scope) 

79 await instance(receive, send) 

80 

81 

82class _AsyncBackend(TypedDict): 

83 backend: str 

84 backend_options: dict[str, Any] 

85 

86 

87class _Upgrade(Exception): 

88 def __init__(self, session: WebSocketTestSession) -> None: 

89 self.session = session 

90 

91 

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 """ 

100 

101 

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 

114 

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 

129 

130 def __exit__(self, *args: Any) -> bool | None: 

131 return self.exit_stack.__exit__(*args) 

132 

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) 

144 

145 # wait for cs.cancel to be called before closing streams 

146 await anyio.sleep_forever() 

147 

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)) 

162 

163 def send(self, message: Message) -> None: 

164 self.portal.call(self._receive_tx.send, message) 

165 

166 def send_text(self, data: str) -> None: 

167 self.send({"type": "websocket.receive", "text": data}) 

168 

169 def send_bytes(self, data: bytes) -> None: 

170 self.send({"type": "websocket.receive", "bytes": data}) 

171 

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")}) 

178 

179 def close(self, code: int = 1000, reason: str | None = None) -> None: 

180 self.send({"type": "websocket.disconnect", "code": code, "reason": reason}) 

181 

182 def receive(self) -> Message: 

183 return self.portal.call(self._send_rx.receive) 

184 

185 def receive_text(self) -> str: 

186 message = self.receive() 

187 self._raise_on_close(message) 

188 return cast(str, message["text"]) 

189 

190 def receive_bytes(self) -> bytes: 

191 message = self.receive() 

192 self._raise_on_close(message) 

193 return cast(bytes, message["bytes"]) 

194 

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) 

203 

204 

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 

222 

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") 

229 

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 

234 

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)] 

240 

241 # Include other request headers. 

242 headers += [(key.lower().encode(), value.encode()) for key, value in request.headers.multi_items()] 

243 

244 scope: dict[str, Any] 

245 

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) 

268 

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 } 

284 

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 

293 

294 async def receive() -> Message: 

295 nonlocal request_complete 

296 

297 if request_complete: 

298 if not response_complete.is_set(): 

299 await response_complete.wait() 

300 return {"type": "http.disconnect"} 

301 

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 

318 

319 request_complete = True 

320 return {"type": "http.request", "body": body_bytes} 

321 

322 async def send(message: Message) -> None: 

323 nonlocal raw_kwargs, response_started, debug_info, body_complete, trailers_expected 

324 

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"] 

353 

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 

361 

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 } 

370 

371 raw_kwargs["stream"] = httpx.ByteStream(raw_kwargs["stream"].read()) 

372 

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 

383 

384 

385class TestClient(httpx.Client): 

386 __test__ = False 

387 task: Future[None] 

388 portal: anyio.from_thread.BlockingPortal | None = None 

389 

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 ) 

429 

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 

437 

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 ) 

478 

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 

499 

500 return session 

501 

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)) 

505 

506 @stack.callback 

507 def reset_portal() -> None: 

508 self.portal = None 

509 

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) 

518 

519 @stack.callback 

520 def wait_shutdown() -> None: 

521 portal.call(self.wait_shutdown) 

522 

523 self.exit_stack = stack.pop_all() 

524 

525 return self 

526 

527 def __exit__(self, *args: Any) -> None: 

528 self.exit_stack.close() 

529 

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) 

536 

537 async def wait_startup(self) -> None: 

538 await self.stream_receive.send({"type": "lifespan.startup"}) 

539 

540 async def receive() -> Any: 

541 message = await self.stream_send.receive() 

542 if message is None: 

543 self.task.result() 

544 return message 

545 

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() 

553 

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 

560 

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()