1from __future__ import annotations
2
3import errno
4import importlib.util
5import os
6import stat
7from email.utils import parsedate
8from typing import Union
9
10import anyio
11import anyio.to_thread
12
13from starlette._utils import get_route_path
14from starlette.datastructures import URL, Headers
15from starlette.exceptions import HTTPException
16from starlette.responses import FileResponse, RedirectResponse, Response
17from starlette.types import Receive, Scope, Send
18from starlette.websockets import WebSocketClose
19
20PathLike = Union[str, "os.PathLike[str]"]
21
22
23class NotModifiedResponse(Response):
24 NOT_MODIFIED_HEADERS = (
25 "cache-control",
26 "content-location",
27 "date",
28 "etag",
29 "expires",
30 "vary",
31 )
32
33 def __init__(self, headers: Headers):
34 super().__init__(
35 status_code=304,
36 headers={name: value for name, value in headers.items() if name in self.NOT_MODIFIED_HEADERS},
37 )
38
39
40class StaticFiles:
41 def __init__(
42 self,
43 *,
44 directory: PathLike | None = None,
45 packages: list[str | tuple[str, str]] | None = None,
46 html: bool = False,
47 check_dir: bool = True,
48 follow_symlink: bool = False,
49 ) -> None:
50 self.directory = directory
51 self.packages = packages
52 self.all_directories = self.get_directories(directory, packages)
53 self.html = html
54 self.config_checked = False
55 self.follow_symlink = follow_symlink
56 if check_dir and directory is not None and not os.path.isdir(directory):
57 raise RuntimeError(f"Directory '{directory}' does not exist")
58
59 def get_directories(
60 self,
61 directory: PathLike | None = None,
62 packages: list[str | tuple[str, str]] | None = None,
63 ) -> list[PathLike]:
64 """
65 Given `directory` and `packages` arguments, return a list of all the
66 directories that should be used for serving static files from.
67 """
68 directories = []
69 if directory is not None:
70 directories.append(directory)
71
72 for package in packages or []:
73 if isinstance(package, tuple):
74 package, statics_dir = package
75 else:
76 statics_dir = "statics"
77 spec = importlib.util.find_spec(package)
78 assert spec is not None, f"Package {package!r} could not be found."
79 assert spec.origin is not None, f"Package {package!r} could not be found."
80 package_directory = os.path.normpath(os.path.join(spec.origin, "..", statics_dir))
81 assert os.path.isdir(package_directory), (
82 f"Directory '{statics_dir!r}' in package {package!r} could not be found."
83 )
84 directories.append(package_directory)
85
86 return directories
87
88 async def __call__(self, scope: Scope, receive: Receive, send: Send) -> None:
89 """
90 The ASGI entry point.
91 """
92 if scope["type"] == "websocket":
93 websocket_close = WebSocketClose()
94 await websocket_close(scope, receive, send)
95 return
96
97 assert scope["type"] == "http"
98
99 if not self.config_checked:
100 await self.check_config()
101 self.config_checked = True
102
103 path = self.get_path(scope)
104 response = await self.get_response(path, scope)
105 await response(scope, receive, send)
106
107 def get_path(self, scope: Scope) -> str:
108 """
109 Given the ASGI scope, return the `path` string to serve up,
110 with OS specific path separators, and any '..', '.' components removed.
111 """
112 route_path = get_route_path(scope)
113 return os.path.normpath(os.path.join(*route_path.split("/")))
114
115 async def get_response(self, path: str, scope: Scope) -> Response:
116 """
117 Returns an HTTP response, given the incoming path, method and request headers.
118 """
119 if scope["method"] not in ("GET", "HEAD"):
120 raise HTTPException(status_code=405)
121
122 try:
123 full_path, stat_result = await anyio.to_thread.run_sync(self.lookup_path, path)
124 except PermissionError:
125 raise HTTPException(status_code=401)
126 except OSError as exc:
127 # Filename is too long, so it can't be a valid static file.
128 if exc.errno == errno.ENAMETOOLONG:
129 raise HTTPException(status_code=404)
130
131 raise exc
132 except ValueError:
133 # Null bytes or other invalid characters in the path.
134 raise HTTPException(status_code=404)
135
136 if stat_result and stat.S_ISREG(stat_result.st_mode):
137 # We have a static file to serve.
138 return self.file_response(full_path, stat_result, scope)
139
140 elif stat_result and stat.S_ISDIR(stat_result.st_mode) and self.html:
141 # We're in HTML mode, and have got a directory URL.
142 # Check if we have 'index.html' file to serve.
143 index_path = os.path.join(path, "index.html")
144 full_path, stat_result = await anyio.to_thread.run_sync(self.lookup_path, index_path)
145 if stat_result is not None and stat.S_ISREG(stat_result.st_mode):
146 if not scope["path"].endswith("/"):
147 # Directory URLs should redirect to always end in "/".
148 url = URL(scope=scope)
149 url = url.replace(path=url.path + "/")
150 return RedirectResponse(url=url)
151 return self.file_response(full_path, stat_result, scope)
152
153 if self.html:
154 # Check for '404.html' if we're in HTML mode.
155 full_path, stat_result = await anyio.to_thread.run_sync(self.lookup_path, "404.html")
156 if stat_result and stat.S_ISREG(stat_result.st_mode):
157 return FileResponse(full_path, stat_result=stat_result, status_code=404)
158 raise HTTPException(status_code=404)
159
160 def lookup_path(self, path: str) -> tuple[str, os.stat_result | None]:
161 # Reject absolute paths so they cannot escape the served directory.
162 if path.startswith(("/", "\\")):
163 return "", None
164 for directory in self.all_directories:
165 joined_path = os.path.join(directory, path)
166 if self.follow_symlink:
167 full_path = os.path.abspath(joined_path)
168 directory = os.path.abspath(directory)
169 else:
170 full_path = os.path.realpath(joined_path)
171 directory = os.path.realpath(directory)
172 if os.path.commonpath([full_path, directory]) != str(directory):
173 # Don't allow misbehaving clients to break out of the static files directory.
174 continue
175 try:
176 return full_path, os.stat(full_path)
177 except (FileNotFoundError, NotADirectoryError):
178 continue
179 return "", None
180
181 def file_response(
182 self,
183 full_path: PathLike,
184 stat_result: os.stat_result,
185 scope: Scope,
186 status_code: int = 200,
187 ) -> Response:
188 request_headers = Headers(scope=scope)
189
190 response = FileResponse(full_path, status_code=status_code, stat_result=stat_result)
191 if self.is_not_modified(response.headers, request_headers):
192 return NotModifiedResponse(response.headers)
193 return response
194
195 async def check_config(self) -> None:
196 """
197 Perform a one-off configuration check that StaticFiles is actually
198 pointed at a directory, so that we can raise loud errors rather than
199 just returning 404 responses.
200 """
201 if self.directory is None:
202 return
203
204 try:
205 stat_result = await anyio.to_thread.run_sync(os.stat, self.directory)
206 except FileNotFoundError:
207 raise RuntimeError(f"StaticFiles directory '{self.directory}' does not exist.")
208 if not (stat.S_ISDIR(stat_result.st_mode) or stat.S_ISLNK(stat_result.st_mode)):
209 raise RuntimeError(f"StaticFiles path '{self.directory}' is not a directory.")
210
211 def is_not_modified(self, response_headers: Headers, request_headers: Headers) -> bool:
212 """
213 Given the request and response headers, return `True` if an HTTP
214 "Not Modified" response could be returned instead.
215 """
216 if if_none_match := request_headers.get("if-none-match"):
217 if if_none_match.strip() == "*":
218 return True
219 # The "etag" header is added by FileResponse, so it's always present.
220 etag = response_headers["etag"]
221 return etag in [tag.strip().removeprefix("W/") for tag in if_none_match.split(",")]
222
223 try:
224 if_modified_since = parsedate(request_headers["if-modified-since"])
225 last_modified = parsedate(response_headers["last-modified"])
226 if if_modified_since is not None and last_modified is not None and if_modified_since >= last_modified:
227 return True
228 except KeyError:
229 pass
230
231 return False