config.py 20 KB

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