Coverage for /pythoncovmergedfiles/medio/medio/usr/local/lib/python3.11/site-packages/starlette/responses.py: 22%

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

337 statements  

1from __future__ import annotations 

2 

3import hashlib 

4import http.cookies 

5import json 

6import os 

7import stat 

8import sys 

9from collections.abc import AsyncIterable, Awaitable, Callable, Iterable, Mapping, Sequence 

10from datetime import datetime 

11from email.utils import format_datetime, formatdate 

12from functools import partial 

13from mimetypes import guess_type 

14from secrets import token_hex 

15from typing import Any, Literal 

16from urllib.parse import quote 

17 

18import anyio 

19import anyio.to_thread 

20 

21from starlette._utils import create_collapsing_task_group 

22from starlette.background import BackgroundTask 

23from starlette.concurrency import iterate_in_threadpool 

24from starlette.datastructures import URL, Headers, MutableHeaders 

25from starlette.requests import ClientDisconnect 

26from starlette.types import Message, Receive, Scope, Send 

27 

28 

29class Response: 

30 media_type = None 

31 charset = "utf-8" 

32 

33 def __init__( 

34 self, 

35 content: Any = None, 

36 status_code: int = 200, 

37 headers: Mapping[str, str] | None = None, 

38 media_type: str | None = None, 

39 background: BackgroundTask | None = None, 

40 ) -> None: 

41 self.status_code = status_code 

42 if media_type is not None: 

43 self.media_type = media_type 

44 self.background = background 

45 self.body = self.render(content) 

46 self.init_headers(headers) 

47 

48 def render(self, content: Any) -> bytes | memoryview: 

49 if content is None: 

50 return b"" 

51 if isinstance(content, bytes | memoryview): 

52 return content 

53 return content.encode(self.charset) # type: ignore 

54 

55 def init_headers(self, headers: Mapping[str, str] | None = None) -> None: 

56 if headers is None: 

57 raw_headers: list[tuple[bytes, bytes]] = [] 

58 populate_content_length = True 

59 populate_content_type = True 

60 else: 

61 raw_headers = [(k.lower().encode("latin-1"), v.encode("latin-1")) for k, v in headers.items()] 

62 keys = [h[0] for h in raw_headers] 

63 populate_content_length = b"content-length" not in keys 

64 populate_content_type = b"content-type" not in keys 

65 

66 body = getattr(self, "body", None) 

67 if ( 

68 body is not None 

69 and populate_content_length 

70 and not (self.status_code < 200 or self.status_code in (204, 304)) 

71 ): 

72 content_length = str(len(body)) 

73 raw_headers.append((b"content-length", content_length.encode("latin-1"))) 

74 

75 content_type = self.media_type 

76 if content_type is not None and populate_content_type: 

77 if content_type.startswith("text/") and "charset=" not in content_type.lower(): 

78 content_type += "; charset=" + self.charset 

79 raw_headers.append((b"content-type", content_type.encode("latin-1"))) 

80 

81 self.raw_headers = raw_headers 

82 

83 @property 

84 def headers(self) -> MutableHeaders: 

85 if not hasattr(self, "_headers"): 

86 self._headers = MutableHeaders(raw=self.raw_headers) 

87 return self._headers 

88 

89 def set_cookie( 

90 self, 

91 key: str, 

92 value: str = "", 

93 max_age: int | None = None, 

94 expires: datetime | str | int | None = None, 

95 path: str | None = "/", 

96 domain: str | None = None, 

97 secure: bool = False, 

98 httponly: bool = False, 

99 samesite: Literal["lax", "strict", "none"] | None = "lax", 

100 partitioned: bool = False, 

101 ) -> None: 

102 cookie: http.cookies.BaseCookie[str] = http.cookies.SimpleCookie() 

103 cookie[key] = value 

104 if max_age is not None: 

105 cookie[key]["max-age"] = max_age 

106 if expires is not None: 

107 if isinstance(expires, datetime): 

108 cookie[key]["expires"] = format_datetime(expires, usegmt=True) 

109 else: 

110 cookie[key]["expires"] = expires 

111 if path is not None: 

112 cookie[key]["path"] = path 

113 if domain is not None: 

114 cookie[key]["domain"] = domain 

115 if secure: 

116 cookie[key]["secure"] = True 

117 if httponly: 

118 cookie[key]["httponly"] = True 

119 if samesite is not None: 

120 assert samesite.lower() in [ 

121 "strict", 

122 "lax", 

123 "none", 

124 ], "samesite must be either 'strict', 'lax' or 'none'" 

125 cookie[key]["samesite"] = samesite 

126 if partitioned: 

127 if sys.version_info < (3, 14): 

128 raise ValueError("Partitioned cookies are only supported in Python 3.14 and above.") # pragma: no cover 

129 cookie[key]["partitioned"] = True # pragma: no cover 

130 

131 cookie_val = cookie.output(header="").strip() 

132 self.raw_headers.append((b"set-cookie", cookie_val.encode("latin-1"))) 

133 

134 def delete_cookie( 

135 self, 

136 key: str, 

137 path: str = "/", 

138 domain: str | None = None, 

139 secure: bool = False, 

140 httponly: bool = False, 

141 samesite: Literal["lax", "strict", "none"] | None = "lax", 

142 partitioned: bool = False, 

143 ) -> None: 

144 self.set_cookie( 

145 key, 

146 max_age=0, 

147 expires=0, 

148 path=path, 

149 domain=domain, 

150 secure=secure, 

151 httponly=httponly, 

152 samesite=samesite, 

153 partitioned=partitioned, 

154 ) 

155 

156 def _wrap_websocket_denial_send(self, send: Send) -> Send: 

157 async def wrapped(message: Message) -> None: 

158 message_type = message["type"] 

159 if message_type in {"http.response.start", "http.response.body"}: # pragma: no branch 

160 message = {**message, "type": "websocket." + message_type} 

161 await send(message) 

162 

163 return wrapped 

164 

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

166 if scope["type"] == "websocket": 

167 send = self._wrap_websocket_denial_send(send) 

168 await send({"type": "http.response.start", "status": self.status_code, "headers": self.raw_headers}) 

169 await send({"type": "http.response.body", "body": self.body}) 

170 

171 if self.background is not None: 

172 await self.background() 

173 

174 

175class HTMLResponse(Response): 

176 media_type = "text/html" 

177 

178 

179class PlainTextResponse(Response): 

180 media_type = "text/plain" 

181 

182 

183class JSONResponse(Response): 

184 media_type = "application/json" 

185 

186 def __init__( 

187 self, 

188 content: Any, 

189 status_code: int = 200, 

190 headers: Mapping[str, str] | None = None, 

191 media_type: str | None = None, 

192 background: BackgroundTask | None = None, 

193 ) -> None: 

194 super().__init__(content, status_code, headers, media_type, background) 

195 

196 def render(self, content: Any) -> bytes: 

197 return json.dumps( 

198 content, 

199 ensure_ascii=False, 

200 allow_nan=False, 

201 indent=None, 

202 separators=(",", ":"), 

203 ).encode("utf-8") 

204 

205 

206class RedirectResponse(Response): 

207 def __init__( 

208 self, 

209 url: str | URL, 

210 status_code: int = 307, 

211 headers: Mapping[str, str] | None = None, 

212 background: BackgroundTask | None = None, 

213 ) -> None: 

214 super().__init__(content=b"", status_code=status_code, headers=headers, background=background) 

215 self.headers["location"] = quote(str(url), safe=":/%#?=@[]!$&'()*+,;") 

216 

217 

218Content = str | bytes | memoryview 

219SyncContentStream = Iterable[Content] 

220AsyncContentStream = AsyncIterable[Content] 

221ContentStream = AsyncContentStream | SyncContentStream 

222 

223 

224class StreamingResponse(Response): 

225 body_iterator: AsyncContentStream 

226 

227 def __init__( 

228 self, 

229 content: ContentStream, 

230 status_code: int = 200, 

231 headers: Mapping[str, str] | None = None, 

232 media_type: str | None = None, 

233 background: BackgroundTask | None = None, 

234 ) -> None: 

235 if isinstance(content, AsyncIterable): 

236 self.body_iterator = content 

237 else: 

238 self.body_iterator = iterate_in_threadpool(content) 

239 self.status_code = status_code 

240 self.media_type = self.media_type if media_type is None else media_type 

241 self.background = background 

242 self.init_headers(headers) 

243 

244 async def listen_for_disconnect(self, receive: Receive) -> None: 

245 while True: 

246 message = await receive() 

247 if message["type"] == "http.disconnect": 

248 break 

249 

250 async def stream_response(self, send: Send) -> None: 

251 await send({"type": "http.response.start", "status": self.status_code, "headers": self.raw_headers}) 

252 async for chunk in self.body_iterator: 

253 if not isinstance(chunk, bytes | memoryview): 

254 chunk = chunk.encode(self.charset) 

255 await send({"type": "http.response.body", "body": chunk, "more_body": True}) 

256 

257 await send({"type": "http.response.body", "body": b"", "more_body": False}) 

258 

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

260 if scope["type"] == "websocket": 

261 send = self._wrap_websocket_denial_send(send) 

262 await self.stream_response(send) 

263 if self.background is not None: 

264 await self.background() 

265 return 

266 

267 spec_version = tuple(map(int, scope.get("asgi", {}).get("spec_version", "2.0").split("."))) 

268 

269 if spec_version >= (2, 4): 

270 try: 

271 await self.stream_response(send) 

272 except OSError: 

273 raise ClientDisconnect() 

274 else: 

275 async with create_collapsing_task_group() as task_group: 

276 

277 async def wrap(func: Callable[[], Awaitable[None]]) -> None: 

278 await func() 

279 task_group.cancel_scope.cancel() 

280 

281 task_group.start_soon(wrap, partial(self.stream_response, send)) 

282 await wrap(partial(self.listen_for_disconnect, receive)) 

283 

284 if self.background is not None: 

285 await self.background() 

286 

287 

288class MalformedRangeHeader(Exception): 

289 def __init__(self, content: str = "Malformed range header.") -> None: 

290 self.content = content 

291 

292 

293class RangeNotSatisfiable(Exception): 

294 def __init__(self, max_size: int) -> None: 

295 self.max_size = max_size 

296 

297 

298class FileResponse(Response): 

299 chunk_size = 64 * 1024 

300 max_ranges = 100 

301 

302 def __init__( 

303 self, 

304 path: str | os.PathLike[str], 

305 status_code: int = 200, 

306 headers: Mapping[str, str] | None = None, 

307 media_type: str | None = None, 

308 background: BackgroundTask | None = None, 

309 filename: str | None = None, 

310 stat_result: os.stat_result | None = None, 

311 content_disposition_type: str = "attachment", 

312 ) -> None: 

313 self.path = path 

314 self.status_code = status_code 

315 self.filename = filename 

316 if media_type is None: 

317 media_type = guess_type(filename or path)[0] or "application/octet-stream" 

318 self.media_type = media_type 

319 self.background = background 

320 self.init_headers(headers) 

321 self.headers.setdefault("accept-ranges", "bytes") 

322 if self.filename is not None: 

323 content_disposition_filename = quote(self.filename) 

324 if content_disposition_filename != self.filename: 

325 content_disposition = f"{content_disposition_type}; filename*=utf-8''{content_disposition_filename}" 

326 else: 

327 content_disposition = f'{content_disposition_type}; filename="{self.filename}"' 

328 self.headers.setdefault("content-disposition", content_disposition) 

329 self.stat_result = stat_result 

330 if stat_result is not None: 

331 self.set_stat_headers(stat_result) 

332 

333 def set_stat_headers(self, stat_result: os.stat_result) -> None: 

334 content_length = str(stat_result.st_size) 

335 last_modified = formatdate(stat_result.st_mtime, usegmt=True) 

336 etag_base = str(stat_result.st_mtime) + "-" + str(stat_result.st_size) 

337 etag = f'"{hashlib.md5(etag_base.encode(), usedforsecurity=False).hexdigest()}"' 

338 

339 self.headers.setdefault("content-length", content_length) 

340 self.headers.setdefault("last-modified", last_modified) 

341 self.headers.setdefault("etag", etag) 

342 

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

344 scope_type = scope["type"] 

345 send_header_only = scope_type == "http" and scope["method"].upper() == "HEAD" 

346 send_pathsend = scope_type == "http" and "http.response.pathsend" in scope.get("extensions", {}) 

347 if scope_type == "websocket": 

348 send = self._wrap_websocket_denial_send(send) 

349 

350 if self.stat_result is None: 

351 try: 

352 stat_result = await anyio.to_thread.run_sync(os.stat, self.path) 

353 self.set_stat_headers(stat_result) 

354 except FileNotFoundError: 

355 raise RuntimeError(f"File at path {self.path} does not exist.") 

356 else: 

357 mode = stat_result.st_mode 

358 if not stat.S_ISREG(mode): 

359 raise RuntimeError(f"File at path {self.path} is not a file.") 

360 else: 

361 stat_result = self.stat_result 

362 

363 headers = Headers(scope=scope) 

364 http_range = headers.get("range") 

365 http_if_range = headers.get("if-range") 

366 

367 if ( 

368 self.status_code != 200 

369 or http_range is None 

370 or (http_if_range is not None and not self._should_use_range(http_if_range)) 

371 ): 

372 await self._handle_simple(send, send_header_only, send_pathsend) 

373 else: 

374 try: 

375 ranges = self._parse_range_header(http_range, stat_result.st_size) 

376 except MalformedRangeHeader as exc: 

377 return await PlainTextResponse(exc.content, status_code=400)(scope, receive, send) 

378 except RangeNotSatisfiable as exc: 

379 response = PlainTextResponse(status_code=416, headers={"Content-Range": f"bytes */{exc.max_size}"}) 

380 return await response(scope, receive, send) 

381 

382 if len(ranges) == 0: 

383 await self._handle_simple(send, send_header_only, send_pathsend) 

384 elif len(ranges) == 1: 

385 start, end = ranges[0] 

386 await self._handle_single_range(send, start, end, stat_result.st_size, send_header_only) 

387 else: 

388 await self._handle_multiple_ranges(send, ranges, stat_result.st_size, send_header_only) 

389 

390 if self.background is not None: 

391 await self.background() 

392 

393 async def _handle_simple(self, send: Send, send_header_only: bool, send_pathsend: bool) -> None: 

394 await send({"type": "http.response.start", "status": self.status_code, "headers": self.raw_headers}) 

395 if send_header_only: 

396 await send({"type": "http.response.body", "body": b"", "more_body": False}) 

397 elif send_pathsend: 

398 await send({"type": "http.response.pathsend", "path": str(self.path)}) 

399 else: 

400 async with await anyio.open_file(self.path, mode="rb") as file: 

401 more_body = True 

402 while more_body: 

403 chunk = await file.read(self.chunk_size) 

404 more_body = len(chunk) == self.chunk_size 

405 await send({"type": "http.response.body", "body": chunk, "more_body": more_body}) 

406 

407 async def _handle_single_range( 

408 self, send: Send, start: int, end: int, file_size: int, send_header_only: bool 

409 ) -> None: 

410 headers = MutableHeaders(raw=list(self.raw_headers)) 

411 headers["content-range"] = f"bytes {start}-{end - 1}/{file_size}" 

412 headers["content-length"] = str(end - start) 

413 await send({"type": "http.response.start", "status": 206, "headers": headers.raw}) 

414 if send_header_only: 

415 await send({"type": "http.response.body", "body": b"", "more_body": False}) 

416 else: 

417 async with await anyio.open_file(self.path, mode="rb") as file: 

418 await file.seek(start) 

419 more_body = True 

420 while more_body: 

421 chunk = await file.read(min(self.chunk_size, end - start)) 

422 start += len(chunk) 

423 more_body = len(chunk) == self.chunk_size and start < end 

424 await send({"type": "http.response.body", "body": chunk, "more_body": more_body}) 

425 

426 async def _handle_multiple_ranges( 

427 self, 

428 send: Send, 

429 ranges: list[tuple[int, int]], 

430 file_size: int, 

431 send_header_only: bool, 

432 ) -> None: 

433 # In firefox and chrome, they use boundary with 95-96 bits entropy (that's roughly 13 bytes). 

434 boundary = token_hex(13) 

435 content_length, header_generator = self.generate_multipart( 

436 ranges, boundary, file_size, self.headers["content-type"] 

437 ) 

438 headers = MutableHeaders(raw=list(self.raw_headers)) 

439 headers["content-type"] = f"multipart/byteranges; boundary={boundary}" 

440 headers["content-length"] = str(content_length) 

441 await send({"type": "http.response.start", "status": 206, "headers": headers.raw}) 

442 if send_header_only: 

443 await send({"type": "http.response.body", "body": b"", "more_body": False}) 

444 else: 

445 async with await anyio.open_file(self.path, mode="rb") as file: 

446 for start, end in ranges: 

447 await send({"type": "http.response.body", "body": header_generator(start, end), "more_body": True}) 

448 await file.seek(start) 

449 while start < end: 

450 chunk = await file.read(min(self.chunk_size, end - start)) 

451 start += len(chunk) 

452 await send({"type": "http.response.body", "body": chunk, "more_body": True}) 

453 await send({"type": "http.response.body", "body": b"\r\n", "more_body": True}) 

454 await send( 

455 { 

456 "type": "http.response.body", 

457 "body": f"--{boundary}--".encode("latin-1"), 

458 "more_body": False, 

459 } 

460 ) 

461 

462 def _should_use_range(self, http_if_range: str) -> bool: 

463 return http_if_range == self.headers["last-modified"] or http_if_range == self.headers["etag"] 

464 

465 @classmethod 

466 def _parse_range_header(cls, http_range: str, file_size: int) -> list[tuple[int, int]]: 

467 ranges: list[tuple[int, int]] = [] 

468 try: 

469 units, range_ = http_range.split("=", 1) 

470 except ValueError: 

471 raise MalformedRangeHeader() 

472 

473 units = units.strip().lower() 

474 

475 if units != "bytes": 

476 raise MalformedRangeHeader("Only support bytes range") 

477 

478 if range_.count(",") + 1 > cls.max_ranges: 

479 return [] 

480 

481 ranges = cls._parse_ranges(range_, file_size) 

482 

483 if len(ranges) == 0: 

484 raise MalformedRangeHeader("Range header: range must be requested") 

485 

486 if any(not (0 <= start < file_size) for start, _ in ranges): 

487 raise RangeNotSatisfiable(file_size) 

488 

489 if any(start >= end for start, end in ranges): 

490 raise MalformedRangeHeader("Range header: start must be less than end") 

491 

492 if len(ranges) == 1: 

493 return ranges 

494 

495 # Merge overlapping ranges 

496 ranges.sort() 

497 result: list[tuple[int, int]] = [ranges[0]] 

498 for start, end in ranges[1:]: 

499 last_start, last_end = result[-1] 

500 if start <= last_end: 

501 result[-1] = (last_start, max(last_end, end)) 

502 else: 

503 result.append((start, end)) 

504 

505 return result 

506 

507 @classmethod 

508 def _parse_ranges(cls, range_: str, file_size: int) -> list[tuple[int, int]]: 

509 ranges: list[tuple[int, int]] = [] 

510 

511 for part in range_.split(","): 

512 part = part.strip() 

513 

514 # If the range is empty or a single dash, we ignore it. 

515 if not part or part == "-": 

516 continue 

517 

518 # If the range is not in the format "start-end", we ignore it. 

519 if "-" not in part: 

520 continue 

521 

522 start_str, end_str = part.split("-", 1) 

523 start_str = start_str.strip() 

524 end_str = end_str.strip() 

525 

526 try: 

527 start = int(start_str) if start_str else max(file_size - int(end_str), 0) 

528 end = int(end_str) + 1 if start_str and end_str and int(end_str) < file_size else file_size 

529 ranges.append((start, end)) 

530 except ValueError: 

531 # If the range is not numeric, we ignore it. 

532 continue 

533 

534 return ranges 

535 

536 def generate_multipart( 

537 self, 

538 ranges: Sequence[tuple[int, int]], 

539 boundary: str, 

540 max_size: int, 

541 content_type: str, 

542 ) -> tuple[int, Callable[[int, int], bytes]]: 

543 r""" 

544 Multipart response headers generator. 

545 

546 ``` 

547 --{boundary}\r\n 

548 Content-Type: {content_type}\r\n 

549 Content-Range: bytes {start}-{end-1}/{max_size}\r\n 

550 \r\n 

551 ..........content...........\r\n 

552 --{boundary}\r\n 

553 Content-Type: {content_type}\r\n 

554 Content-Range: bytes {start}-{end-1}/{max_size}\r\n 

555 \r\n 

556 ..........content...........\r\n 

557 --{boundary}-- 

558 ``` 

559 """ 

560 boundary_len = len(boundary) 

561 static_header_part_len = 49 + boundary_len + len(content_type) + len(str(max_size)) 

562 content_length = sum( 

563 (len(str(start)) + len(str(end - 1)) + static_header_part_len) # Headers 

564 + (end - start) # Content 

565 for start, end in ranges 

566 ) + ( 

567 4 + boundary_len # --boundary-- 

568 ) 

569 return ( 

570 content_length, 

571 lambda start, end: ( 

572 f"""\ 

573--{boundary}\r 

574Content-Type: {content_type}\r 

575Content-Range: bytes {start}-{end - 1}/{max_size}\r 

576\r 

577""" 

578 ).encode("latin-1"), 

579 )