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
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 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
18import anyio
19import anyio.to_thread
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
29class Response:
30 media_type = None
31 charset = "utf-8"
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)
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
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
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")))
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")))
81 self.raw_headers = raw_headers
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
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
131 cookie_val = cookie.output(header="").strip()
132 self.raw_headers.append((b"set-cookie", cookie_val.encode("latin-1")))
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 )
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)
163 return wrapped
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})
171 if self.background is not None:
172 await self.background()
175class HTMLResponse(Response):
176 media_type = "text/html"
179class PlainTextResponse(Response):
180 media_type = "text/plain"
183class JSONResponse(Response):
184 media_type = "application/json"
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)
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")
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=":/%#?=@[]!$&'()*+,;")
218Content = str | bytes | memoryview
219SyncContentStream = Iterable[Content]
220AsyncContentStream = AsyncIterable[Content]
221ContentStream = AsyncContentStream | SyncContentStream
224class StreamingResponse(Response):
225 body_iterator: AsyncContentStream
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)
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
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})
257 await send({"type": "http.response.body", "body": b"", "more_body": False})
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
267 spec_version = tuple(map(int, scope.get("asgi", {}).get("spec_version", "2.0").split(".")))
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:
277 async def wrap(func: Callable[[], Awaitable[None]]) -> None:
278 await func()
279 task_group.cancel_scope.cancel()
281 task_group.start_soon(wrap, partial(self.stream_response, send))
282 await wrap(partial(self.listen_for_disconnect, receive))
284 if self.background is not None:
285 await self.background()
288class MalformedRangeHeader(Exception):
289 def __init__(self, content: str = "Malformed range header.") -> None:
290 self.content = content
293class RangeNotSatisfiable(Exception):
294 def __init__(self, max_size: int) -> None:
295 self.max_size = max_size
298class FileResponse(Response):
299 chunk_size = 64 * 1024
300 max_ranges = 100
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)
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()}"'
339 self.headers.setdefault("content-length", content_length)
340 self.headers.setdefault("last-modified", last_modified)
341 self.headers.setdefault("etag", etag)
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)
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
363 headers = Headers(scope=scope)
364 http_range = headers.get("range")
365 http_if_range = headers.get("if-range")
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)
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)
390 if self.background is not None:
391 await self.background()
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})
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})
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 )
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"]
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()
473 units = units.strip().lower()
475 if units != "bytes":
476 raise MalformedRangeHeader("Only support bytes range")
478 if range_.count(",") + 1 > cls.max_ranges:
479 return []
481 ranges = cls._parse_ranges(range_, file_size)
483 if len(ranges) == 0:
484 raise MalformedRangeHeader("Range header: range must be requested")
486 if any(not (0 <= start < file_size) for start, _ in ranges):
487 raise RangeNotSatisfiable(file_size)
489 if any(start >= end for start, end in ranges):
490 raise MalformedRangeHeader("Range header: start must be less than end")
492 if len(ranges) == 1:
493 return ranges
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))
505 return result
507 @classmethod
508 def _parse_ranges(cls, range_: str, file_size: int) -> list[tuple[int, int]]:
509 ranges: list[tuple[int, int]] = []
511 for part in range_.split(","):
512 part = part.strip()
514 # If the range is empty or a single dash, we ignore it.
515 if not part or part == "-":
516 continue
518 # If the range is not in the format "start-end", we ignore it.
519 if "-" not in part:
520 continue
522 start_str, end_str = part.split("-", 1)
523 start_str = start_str.strip()
524 end_str = end_str.strip()
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
534 return ranges
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.
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 )