config.py 20 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218219220221222223224225226227228229230231232233234235236237238239240241242243244245246247248249250251252253254255256257258259260261262263264265266267268269270271272273274275276277278279280281282283284285286287288289290291292293294295296297298299300301302303304305306307308309310311312313314315316317318319320321322323324325326327328329330331332333334335336337338339340341342343344345346347348349350351352353354355356357358359360361362363364365366367368369370371372373374375376377378379380381382383384385386387388389390391392393394395396397398399400401402403404405406407408409410411412413414415416417418419420421422423424425426427428429430431432433434435436437438439440441442443444445446447448449450451452453454455456457458459460461462463464465466467468469470471472473474475476477478479480481482483484485486487488489490491492493494495496497498499500501502503504505506507508509510511512513514515516517518519520521522523524525526527528529
  1. from __future__ import annotations
  2. import asyncio
  3. import inspect
  4. import json
  5. import logging
  6. import logging.config
  7. import os
  8. import socket
  9. import ssl
  10. import sys
  11. from configparser import RawConfigParser
  12. from pathlib import Path
  13. from typing import IO, Any, Awaitable, Callable, Literal
  14. import click
  15. from uvicorn._types import ASGIApplication
  16. from uvicorn.importer import ImportFromStringError, import_from_string
  17. from uvicorn.logging import TRACE_LOG_LEVEL
  18. from uvicorn.middleware.asgi2 import ASGI2Middleware
  19. from uvicorn.middleware.message_logger import MessageLoggerMiddleware
  20. from uvicorn.middleware.proxy_headers import ProxyHeadersMiddleware
  21. from uvicorn.middleware.wsgi import WSGIMiddleware
  22. HTTPProtocolType = Literal["auto", "h11", "httptools"]
  23. WSProtocolType = Literal["auto", "none", "websockets", "wsproto"]
  24. LifespanType = Literal["auto", "on", "off"]
  25. LoopSetupType = Literal["none", "auto", "asyncio", "uvloop"]
  26. InterfaceType = Literal["auto", "asgi3", "asgi2", "wsgi"]
  27. LOG_LEVELS: dict[str, int] = {
  28. "critical": logging.CRITICAL,
  29. "error": logging.ERROR,
  30. "warning": logging.WARNING,
  31. "info": logging.INFO,
  32. "debug": logging.DEBUG,
  33. "trace": TRACE_LOG_LEVEL,
  34. }
  35. HTTP_PROTOCOLS: dict[HTTPProtocolType, str] = {
  36. "auto": "uvicorn.protocols.http.auto:AutoHTTPProtocol",
  37. "h11": "uvicorn.protocols.http.h11_impl:H11Protocol",
  38. "httptools": "uvicorn.protocols.http.httptools_impl:HttpToolsProtocol",
  39. }
  40. WS_PROTOCOLS: dict[WSProtocolType, str | None] = {
  41. "auto": "uvicorn.protocols.websockets.auto:AutoWebSocketsProtocol",
  42. "none": None,
  43. "websockets": "uvicorn.protocols.websockets.websockets_impl:WebSocketProtocol",
  44. "wsproto": "uvicorn.protocols.websockets.wsproto_impl:WSProtocol",
  45. }
  46. LIFESPAN: dict[LifespanType, str] = {
  47. "auto": "uvicorn.lifespan.on:LifespanOn",
  48. "on": "uvicorn.lifespan.on:LifespanOn",
  49. "off": "uvicorn.lifespan.off:LifespanOff",
  50. }
  51. LOOP_SETUPS: dict[LoopSetupType, str | None] = {
  52. "none": None,
  53. "auto": "uvicorn.loops.auto:auto_loop_setup",
  54. "asyncio": "uvicorn.loops.asyncio:asyncio_setup",
  55. "uvloop": "uvicorn.loops.uvloop:uvloop_setup",
  56. }
  57. INTERFACES: list[InterfaceType] = ["auto", "asgi3", "asgi2", "wsgi"]
  58. SSL_PROTOCOL_VERSION: int = ssl.PROTOCOL_TLS_SERVER
  59. LOGGING_CONFIG: dict[str, Any] = {
  60. "version": 1,
  61. "disable_existing_loggers": False,
  62. "formatters": {
  63. "default": {
  64. "()": "uvicorn.logging.DefaultFormatter",
  65. "fmt": "%(levelprefix)s %(message)s",
  66. "use_colors": None,
  67. },
  68. "access": {
  69. "()": "uvicorn.logging.AccessFormatter",
  70. "fmt": '%(levelprefix)s %(client_addr)s - "%(request_line)s" %(status_code)s', # noqa: E501
  71. },
  72. },
  73. "handlers": {
  74. "default": {
  75. "formatter": "default",
  76. "class": "logging.StreamHandler",
  77. "stream": "ext://sys.stderr",
  78. },
  79. "access": {
  80. "formatter": "access",
  81. "class": "logging.StreamHandler",
  82. "stream": "ext://sys.stdout",
  83. },
  84. },
  85. "loggers": {
  86. "uvicorn": {"handlers": ["default"], "level": "INFO", "propagate": False},
  87. "uvicorn.error": {"level": "INFO"},
  88. "uvicorn.access": {"handlers": ["access"], "level": "INFO", "propagate": False},
  89. },
  90. }
  91. logger = logging.getLogger("uvicorn.error")
  92. def create_ssl_context(
  93. certfile: str | os.PathLike[str],
  94. keyfile: str | os.PathLike[str] | None,
  95. password: str | None,
  96. ssl_version: int,
  97. cert_reqs: int,
  98. ca_certs: str | os.PathLike[str] | None,
  99. ciphers: str | None,
  100. ) -> ssl.SSLContext:
  101. ctx = ssl.SSLContext(ssl_version)
  102. get_password = (lambda: password) if password else None
  103. ctx.load_cert_chain(certfile, keyfile, get_password)
  104. ctx.verify_mode = ssl.VerifyMode(cert_reqs)
  105. if ca_certs:
  106. ctx.load_verify_locations(ca_certs)
  107. if ciphers:
  108. ctx.set_ciphers(ciphers)
  109. return ctx
  110. def is_dir(path: Path) -> bool:
  111. try:
  112. if not path.is_absolute():
  113. path = path.resolve()
  114. return path.is_dir()
  115. except OSError: # pragma: full coverage
  116. return False
  117. def resolve_reload_patterns(patterns_list: list[str], directories_list: list[str]) -> tuple[list[str], list[Path]]:
  118. directories: list[Path] = list(set(map(Path, directories_list.copy())))
  119. patterns: list[str] = patterns_list.copy()
  120. current_working_directory = Path.cwd()
  121. for pattern in patterns_list:
  122. # Special case for the .* pattern, otherwise this would only match
  123. # hidden directories which is probably undesired
  124. if pattern == ".*":
  125. continue
  126. patterns.append(pattern)
  127. if is_dir(Path(pattern)):
  128. directories.append(Path(pattern))
  129. else:
  130. for match in current_working_directory.glob(pattern):
  131. if is_dir(match):
  132. directories.append(match)
  133. directories = list(set(directories))
  134. directories = list(map(Path, directories))
  135. directories = list(map(lambda x: x.resolve(), directories))
  136. directories = list({reload_path for reload_path in directories if is_dir(reload_path)})
  137. children = []
  138. for j in range(len(directories)):
  139. for k in range(j + 1, len(directories)): # pragma: full coverage
  140. if directories[j] in directories[k].parents:
  141. children.append(directories[k])
  142. elif directories[k] in directories[j].parents:
  143. children.append(directories[j])
  144. directories = list(set(directories).difference(set(children)))
  145. return list(set(patterns)), directories
  146. def _normalize_dirs(dirs: list[str] | str | None) -> list[str]:
  147. if dirs is None:
  148. return []
  149. if isinstance(dirs, str):
  150. return [dirs]
  151. return list(set(dirs))
  152. class Config:
  153. def __init__(
  154. self,
  155. app: ASGIApplication | Callable[..., Any] | str,
  156. host: str = "127.0.0.1",
  157. port: int = 8000,
  158. uds: str | None = None,
  159. fd: int | None = None,
  160. loop: LoopSetupType = "auto",
  161. http: type[asyncio.Protocol] | HTTPProtocolType = "auto",
  162. ws: type[asyncio.Protocol] | WSProtocolType = "auto",
  163. ws_max_size: int = 16 * 1024 * 1024,
  164. ws_max_queue: int = 32,
  165. ws_ping_interval: float | None = 20.0,
  166. ws_ping_timeout: float | None = 20.0,
  167. ws_per_message_deflate: bool = True,
  168. lifespan: LifespanType = "auto",
  169. env_file: str | os.PathLike[str] | None = None,
  170. log_config: dict[str, Any] | str | RawConfigParser | IO[Any] | None = LOGGING_CONFIG,
  171. log_level: str | int | None = None,
  172. access_log: bool = True,
  173. use_colors: bool | None = None,
  174. interface: InterfaceType = "auto",
  175. reload: bool = False,
  176. reload_dirs: list[str] | str | None = None,
  177. reload_delay: float = 0.25,
  178. reload_includes: list[str] | str | None = None,
  179. reload_excludes: list[str] | str | None = None,
  180. workers: int | None = None,
  181. proxy_headers: bool = True,
  182. server_header: bool = True,
  183. date_header: bool = True,
  184. forwarded_allow_ips: list[str] | str | None = None,
  185. root_path: str = "",
  186. limit_concurrency: int | None = None,
  187. limit_max_requests: int | None = None,
  188. backlog: int = 2048,
  189. timeout_keep_alive: int = 5,
  190. timeout_notify: int = 30,
  191. timeout_graceful_shutdown: int | None = None,
  192. callback_notify: Callable[..., Awaitable[None]] | None = None,
  193. ssl_keyfile: str | os.PathLike[str] | None = None,
  194. ssl_certfile: str | os.PathLike[str] | None = None,
  195. ssl_keyfile_password: str | None = None,
  196. ssl_version: int = SSL_PROTOCOL_VERSION,
  197. ssl_cert_reqs: int = ssl.CERT_NONE,
  198. ssl_ca_certs: str | None = None,
  199. ssl_ciphers: str = "TLSv1",
  200. headers: list[tuple[str, str]] | None = None,
  201. factory: bool = False,
  202. h11_max_incomplete_event_size: int | None = None,
  203. ):
  204. self.app = app
  205. self.host = host
  206. self.port = port
  207. self.uds = uds
  208. self.fd = fd
  209. self.loop = loop
  210. self.http = http
  211. self.ws = ws
  212. self.ws_max_size = ws_max_size
  213. self.ws_max_queue = ws_max_queue
  214. self.ws_ping_interval = ws_ping_interval
  215. self.ws_ping_timeout = ws_ping_timeout
  216. self.ws_per_message_deflate = ws_per_message_deflate
  217. self.lifespan = lifespan
  218. self.log_config = log_config
  219. self.log_level = log_level
  220. self.access_log = access_log
  221. self.use_colors = use_colors
  222. self.interface = interface
  223. self.reload = reload
  224. self.reload_delay = reload_delay
  225. self.workers = workers or 1
  226. self.proxy_headers = proxy_headers
  227. self.server_header = server_header
  228. self.date_header = date_header
  229. self.root_path = root_path
  230. self.limit_concurrency = limit_concurrency
  231. self.limit_max_requests = limit_max_requests
  232. self.backlog = backlog
  233. self.timeout_keep_alive = timeout_keep_alive
  234. self.timeout_notify = timeout_notify
  235. self.timeout_graceful_shutdown = timeout_graceful_shutdown
  236. self.callback_notify = callback_notify
  237. self.ssl_keyfile = ssl_keyfile
  238. self.ssl_certfile = ssl_certfile
  239. self.ssl_keyfile_password = ssl_keyfile_password
  240. self.ssl_version = ssl_version
  241. self.ssl_cert_reqs = ssl_cert_reqs
  242. self.ssl_ca_certs = ssl_ca_certs
  243. self.ssl_ciphers = ssl_ciphers
  244. self.headers: list[tuple[str, str]] = headers or []
  245. self.encoded_headers: list[tuple[bytes, bytes]] = []
  246. self.factory = factory
  247. self.h11_max_incomplete_event_size = h11_max_incomplete_event_size
  248. self.loaded = False
  249. self.configure_logging()
  250. self.reload_dirs: list[Path] = []
  251. self.reload_dirs_excludes: list[Path] = []
  252. self.reload_includes: list[str] = []
  253. self.reload_excludes: list[str] = []
  254. if (reload_dirs or reload_includes or reload_excludes) and not self.should_reload:
  255. logger.warning(
  256. "Current configuration will not reload as not all conditions are met, " "please refer to documentation."
  257. )
  258. if self.should_reload:
  259. reload_dirs = _normalize_dirs(reload_dirs)
  260. reload_includes = _normalize_dirs(reload_includes)
  261. reload_excludes = _normalize_dirs(reload_excludes)
  262. self.reload_includes, self.reload_dirs = resolve_reload_patterns(reload_includes, reload_dirs)
  263. self.reload_excludes, self.reload_dirs_excludes = resolve_reload_patterns(reload_excludes, [])
  264. reload_dirs_tmp = self.reload_dirs.copy()
  265. for directory in self.reload_dirs_excludes:
  266. for reload_directory in reload_dirs_tmp:
  267. if directory == reload_directory or directory in reload_directory.parents:
  268. try:
  269. self.reload_dirs.remove(reload_directory)
  270. except ValueError: # pragma: full coverage
  271. pass
  272. for pattern in self.reload_excludes:
  273. if pattern in self.reload_includes:
  274. self.reload_includes.remove(pattern) # pragma: full coverage
  275. if not self.reload_dirs:
  276. if reload_dirs:
  277. logger.warning(
  278. "Provided reload directories %s did not contain valid "
  279. + "directories, watching current working directory.",
  280. reload_dirs,
  281. )
  282. self.reload_dirs = [Path(os.getcwd())]
  283. logger.info(
  284. "Will watch for changes in these directories: %s",
  285. sorted(list(map(str, self.reload_dirs))),
  286. )
  287. if env_file is not None:
  288. from dotenv import load_dotenv
  289. logger.info("Loading environment from '%s'", env_file)
  290. load_dotenv(dotenv_path=env_file)
  291. if workers is None and "WEB_CONCURRENCY" in os.environ:
  292. self.workers = int(os.environ["WEB_CONCURRENCY"])
  293. self.forwarded_allow_ips: list[str] | str
  294. if forwarded_allow_ips is None:
  295. self.forwarded_allow_ips = os.environ.get("FORWARDED_ALLOW_IPS", "127.0.0.1")
  296. else:
  297. self.forwarded_allow_ips = forwarded_allow_ips # pragma: full coverage
  298. if self.reload and self.workers > 1:
  299. logger.warning('"workers" flag is ignored when reloading is enabled.')
  300. @property
  301. def asgi_version(self) -> Literal["2.0", "3.0"]:
  302. mapping: dict[str, Literal["2.0", "3.0"]] = {
  303. "asgi2": "2.0",
  304. "asgi3": "3.0",
  305. "wsgi": "3.0",
  306. }
  307. return mapping[self.interface]
  308. @property
  309. def is_ssl(self) -> bool:
  310. return bool(self.ssl_keyfile or self.ssl_certfile)
  311. @property
  312. def use_subprocess(self) -> bool:
  313. return bool(self.reload or self.workers > 1)
  314. def configure_logging(self) -> None:
  315. logging.addLevelName(TRACE_LOG_LEVEL, "TRACE")
  316. if self.log_config is not None:
  317. if isinstance(self.log_config, dict):
  318. if self.use_colors in (True, False):
  319. self.log_config["formatters"]["default"]["use_colors"] = self.use_colors
  320. self.log_config["formatters"]["access"]["use_colors"] = self.use_colors
  321. logging.config.dictConfig(self.log_config)
  322. elif isinstance(self.log_config, str) and self.log_config.endswith(".json"):
  323. with open(self.log_config) as file:
  324. loaded_config = json.load(file)
  325. logging.config.dictConfig(loaded_config)
  326. elif isinstance(self.log_config, str) and self.log_config.endswith((".yaml", ".yml")):
  327. # Install the PyYAML package or the uvicorn[standard] optional
  328. # dependencies to enable this functionality.
  329. import yaml
  330. with open(self.log_config) as file:
  331. loaded_config = yaml.safe_load(file)
  332. logging.config.dictConfig(loaded_config)
  333. else:
  334. # See the note about fileConfig() here:
  335. # https://docs.python.org/3/library/logging.config.html#configuration-file-format
  336. logging.config.fileConfig(self.log_config, disable_existing_loggers=False)
  337. if self.log_level is not None:
  338. if isinstance(self.log_level, str):
  339. log_level = LOG_LEVELS[self.log_level]
  340. else:
  341. log_level = self.log_level
  342. logging.getLogger("uvicorn.error").setLevel(log_level)
  343. logging.getLogger("uvicorn.access").setLevel(log_level)
  344. logging.getLogger("uvicorn.asgi").setLevel(log_level)
  345. if self.access_log is False:
  346. logging.getLogger("uvicorn.access").handlers = []
  347. logging.getLogger("uvicorn.access").propagate = False
  348. def load(self) -> None:
  349. assert not self.loaded
  350. if self.is_ssl:
  351. assert self.ssl_certfile
  352. self.ssl: ssl.SSLContext | None = create_ssl_context(
  353. keyfile=self.ssl_keyfile,
  354. certfile=self.ssl_certfile,
  355. password=self.ssl_keyfile_password,
  356. ssl_version=self.ssl_version,
  357. cert_reqs=self.ssl_cert_reqs,
  358. ca_certs=self.ssl_ca_certs,
  359. ciphers=self.ssl_ciphers,
  360. )
  361. else:
  362. self.ssl = None
  363. encoded_headers = [(key.lower().encode("latin1"), value.encode("latin1")) for key, value in self.headers]
  364. self.encoded_headers = (
  365. [(b"server", b"uvicorn")] + encoded_headers
  366. if b"server" not in dict(encoded_headers) and self.server_header
  367. else encoded_headers
  368. )
  369. if isinstance(self.http, str):
  370. http_protocol_class = import_from_string(HTTP_PROTOCOLS[self.http])
  371. self.http_protocol_class: type[asyncio.Protocol] = http_protocol_class
  372. else:
  373. self.http_protocol_class = self.http
  374. if isinstance(self.ws, str):
  375. ws_protocol_class = import_from_string(WS_PROTOCOLS[self.ws])
  376. self.ws_protocol_class: type[asyncio.Protocol] | None = ws_protocol_class
  377. else:
  378. self.ws_protocol_class = self.ws
  379. self.lifespan_class = import_from_string(LIFESPAN[self.lifespan])
  380. try:
  381. self.loaded_app = import_from_string(self.app)
  382. except ImportFromStringError as exc:
  383. logger.error("Error loading ASGI app. %s" % exc)
  384. sys.exit(1)
  385. try:
  386. self.loaded_app = self.loaded_app()
  387. except TypeError as exc:
  388. if self.factory:
  389. logger.error("Error loading ASGI app factory: %s", exc)
  390. sys.exit(1)
  391. else:
  392. if not self.factory:
  393. logger.warning(
  394. "ASGI app factory detected. Using it, " "but please consider setting the --factory flag explicitly."
  395. )
  396. if self.interface == "auto":
  397. if inspect.isclass(self.loaded_app):
  398. use_asgi_3 = hasattr(self.loaded_app, "__await__")
  399. elif inspect.isfunction(self.loaded_app):
  400. use_asgi_3 = asyncio.iscoroutinefunction(self.loaded_app)
  401. else:
  402. call = getattr(self.loaded_app, "__call__", None)
  403. use_asgi_3 = asyncio.iscoroutinefunction(call)
  404. self.interface = "asgi3" if use_asgi_3 else "asgi2"
  405. if self.interface == "wsgi":
  406. self.loaded_app = WSGIMiddleware(self.loaded_app)
  407. self.ws_protocol_class = None
  408. elif self.interface == "asgi2":
  409. self.loaded_app = ASGI2Middleware(self.loaded_app)
  410. if logger.getEffectiveLevel() <= TRACE_LOG_LEVEL:
  411. self.loaded_app = MessageLoggerMiddleware(self.loaded_app)
  412. if self.proxy_headers:
  413. self.loaded_app = ProxyHeadersMiddleware(self.loaded_app, trusted_hosts=self.forwarded_allow_ips)
  414. self.loaded = True
  415. def setup_event_loop(self) -> None:
  416. loop_setup: Callable | None = import_from_string(LOOP_SETUPS[self.loop])
  417. if loop_setup is not None:
  418. loop_setup(use_subprocess=self.use_subprocess)
  419. def bind_socket(self) -> socket.socket:
  420. logger_args: list[str | int]
  421. if self.uds: # pragma: py-win32
  422. path = self.uds
  423. sock = socket.socket(socket.AF_UNIX, socket.SOCK_STREAM)
  424. try:
  425. sock.bind(path)
  426. uds_perms = 0o666
  427. os.chmod(self.uds, uds_perms)
  428. except OSError as exc: # pragma: full coverage
  429. logger.error(exc)
  430. sys.exit(1)
  431. message = "Uvicorn running on unix socket %s (Press CTRL+C to quit)"
  432. sock_name_format = "%s"
  433. color_message = "Uvicorn running on " + click.style(sock_name_format, bold=True) + " (Press CTRL+C to quit)"
  434. logger_args = [self.uds]
  435. elif self.fd: # pragma: py-win32
  436. sock = socket.fromfd(self.fd, socket.AF_UNIX, socket.SOCK_STREAM)
  437. message = "Uvicorn running on socket %s (Press CTRL+C to quit)"
  438. fd_name_format = "%s"
  439. color_message = "Uvicorn running on " + click.style(fd_name_format, bold=True) + " (Press CTRL+C to quit)"
  440. logger_args = [sock.getsockname()]
  441. else:
  442. family = socket.AF_INET
  443. addr_format = "%s://%s:%d"
  444. if self.host and ":" in self.host: # pragma: full coverage
  445. # It's an IPv6 address.
  446. family = socket.AF_INET6
  447. addr_format = "%s://[%s]:%d"
  448. sock = socket.socket(family=family)
  449. sock.setsockopt(socket.SOL_SOCKET, socket.SO_REUSEADDR, 1)
  450. try:
  451. sock.bind((self.host, self.port))
  452. except OSError as exc: # pragma: full coverage
  453. logger.error(exc)
  454. sys.exit(1)
  455. message = f"Uvicorn running on {addr_format} (Press CTRL+C to quit)"
  456. color_message = "Uvicorn running on " + click.style(addr_format, bold=True) + " (Press CTRL+C to quit)"
  457. protocol_name = "https" if self.is_ssl else "http"
  458. logger_args = [protocol_name, self.host, sock.getsockname()[1]]
  459. logger.info(message, *logger_args, extra={"color_message": color_message})
  460. sock.set_inheritable(True)
  461. return sock
  462. @property
  463. def should_reload(self) -> bool:
  464. return isinstance(self.app, str) and self.reload