server.py 13 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218219220221222223224225226227228229230231232233234235236237238239240241242243244245246247248249250251252253254255256257258259260261262263264265266267268269270271272273274275276277278279280281282283284285286287288289290291292293294295296297298299300301302303304305306307308309310311312313314315316317318319320321322323324325326327328329330331332333334335336337338339340341342343344345346347
  1. from __future__ import annotations
  2. import asyncio
  3. import contextlib
  4. import functools
  5. import logging
  6. import os
  7. import platform
  8. import random
  9. import signal
  10. import socket
  11. import sys
  12. import threading
  13. import time
  14. from collections.abc import Generator, Sequence
  15. from email.utils import formatdate
  16. from types import FrameType
  17. from typing import TYPE_CHECKING, TypeAlias
  18. from uvicorn._ansi import style
  19. from uvicorn._compat import asyncio_run
  20. from uvicorn.config import STARTUP_FAILURE, Config
  21. if TYPE_CHECKING:
  22. from uvicorn.protocols.http.h11_impl import H11Protocol
  23. from uvicorn.protocols.http.httptools_impl import HttpToolsProtocol
  24. from uvicorn.protocols.http.zttp_impl import ZttpProtocol
  25. from uvicorn.protocols.websockets.websockets_impl import WebSocketProtocol
  26. from uvicorn.protocols.websockets.websockets_sansio_impl import WebSocketsSansIOProtocol
  27. from uvicorn.protocols.websockets.wsproto_impl import WSProtocol
  28. Protocols: TypeAlias = (
  29. H11Protocol | HttpToolsProtocol | ZttpProtocol | WSProtocol | WebSocketProtocol | WebSocketsSansIOProtocol
  30. )
  31. HANDLED_SIGNALS = (
  32. signal.SIGINT, # Unix signal 2. Sent by Ctrl+C.
  33. signal.SIGTERM, # Unix signal 15. Sent by `kill <pid>`.
  34. )
  35. if sys.platform == "win32": # pragma: py-not-win32
  36. HANDLED_SIGNALS += (signal.SIGBREAK,) # Windows signal 21. Sent by Ctrl+Break.
  37. logger = logging.getLogger("uvicorn.error")
  38. class ServerState:
  39. """
  40. Shared servers state that is available between all protocol instances.
  41. """
  42. def __init__(self) -> None:
  43. self.total_requests = 0
  44. self.connections: set[Protocols] = set()
  45. self.tasks: set[asyncio.Task[None]] = set()
  46. self.default_headers: list[tuple[bytes, bytes]] = []
  47. class Server:
  48. def __init__(self, config: Config) -> None:
  49. self.config = config
  50. self.server_state = ServerState()
  51. self.started = False
  52. self.should_exit = False
  53. self.force_exit = False
  54. self.last_notified = 0.0
  55. self._captured_signals: list[int] = []
  56. @functools.cached_property
  57. def limit_max_requests(self) -> int | None:
  58. if self.config.limit_max_requests is None:
  59. return None
  60. return self.config.limit_max_requests + random.randint(0, self.config.limit_max_requests_jitter)
  61. def run(self, sockets: list[socket.socket] | None = None) -> None:
  62. return asyncio_run(self.serve(sockets=sockets), loop_factory=self.config.get_loop_factory())
  63. async def serve(self, sockets: list[socket.socket] | None = None) -> None:
  64. with self.capture_signals():
  65. await self._serve(sockets)
  66. async def _serve(self, sockets: list[socket.socket] | None = None) -> None:
  67. process_id = os.getpid()
  68. config = self.config
  69. if not config.loaded:
  70. config.load()
  71. self.lifespan = config.lifespan_class(config)
  72. message = "Started server process [%d]"
  73. color_message = "Started server process [" + style("%d", fg="cyan") + "]"
  74. logger.info(message, process_id, extra={"color_message": color_message})
  75. await self.startup(sockets=sockets)
  76. if not self.should_exit:
  77. await self.main_loop()
  78. if self.started:
  79. await self.shutdown(sockets=sockets)
  80. message = "Finished server process [%d]"
  81. color_message = "Finished server process [" + style("%d", fg="cyan") + "]"
  82. logger.info(message, process_id, extra={"color_message": color_message})
  83. async def startup(self, sockets: list[socket.socket] | None = None) -> None:
  84. await self.lifespan.startup()
  85. if self.lifespan.should_exit:
  86. sys.exit(STARTUP_FAILURE)
  87. config = self.config
  88. def create_protocol(
  89. _loop: asyncio.AbstractEventLoop | None = None,
  90. ) -> asyncio.Protocol:
  91. return config.http_protocol_class( # type: ignore[call-arg]
  92. config=config,
  93. server_state=self.server_state,
  94. app_state=self.lifespan.state,
  95. _loop=_loop,
  96. )
  97. loop = asyncio.get_running_loop()
  98. listeners: Sequence[socket.SocketType]
  99. if sockets is not None: # pragma: full coverage
  100. # Explicitly passed a list of open sockets.
  101. # We use this when the server is run from a Gunicorn worker.
  102. def _share_socket(
  103. sock: socket.SocketType,
  104. ) -> socket.SocketType: # pragma py-not-win32
  105. # Windows requires the socket be explicitly shared across
  106. # multiple workers (processes).
  107. from socket import fromshare # type: ignore[attr-defined]
  108. sock_data = sock.share(os.getpid()) # type: ignore[attr-defined]
  109. return fromshare(sock_data)
  110. self.servers: list[asyncio.base_events.Server] = []
  111. for sock in sockets:
  112. is_windows = platform.system() == "Windows"
  113. if config.workers > 1 and is_windows: # pragma: py-not-win32
  114. sock = _share_socket(sock) # type: ignore[assignment]
  115. server = await loop.create_server(create_protocol, sock=sock, ssl=config.ssl, backlog=config.backlog)
  116. self.servers.append(server)
  117. listeners = sockets
  118. elif config.fd is not None: # pragma: py-win32
  119. # Use an existing socket, from a file descriptor.
  120. sock = socket.fromfd(config.fd, socket.AF_UNIX, socket.SOCK_STREAM)
  121. server = await loop.create_server(create_protocol, sock=sock, ssl=config.ssl, backlog=config.backlog)
  122. assert server.sockets is not None # mypy
  123. listeners = server.sockets
  124. self.servers = [server]
  125. elif config.uds is not None: # pragma: py-win32
  126. # Create a socket using UNIX domain socket.
  127. uds_perms = 0o666
  128. if os.path.exists(config.uds):
  129. uds_perms = os.stat(config.uds).st_mode # pragma: full coverage
  130. server = await loop.create_unix_server(
  131. create_protocol, path=config.uds, ssl=config.ssl, backlog=config.backlog
  132. )
  133. os.chmod(config.uds, uds_perms)
  134. assert server.sockets is not None # mypy
  135. listeners = server.sockets
  136. self.servers = [server]
  137. else:
  138. # Standard case. Create a socket from a host/port pair.
  139. try:
  140. server = await loop.create_server(
  141. create_protocol,
  142. host=config.host,
  143. port=config.port,
  144. ssl=config.ssl,
  145. backlog=config.backlog,
  146. )
  147. except OSError as exc:
  148. logger.error(exc)
  149. await self.lifespan.shutdown()
  150. sys.exit(STARTUP_FAILURE)
  151. assert server.sockets is not None
  152. listeners = server.sockets
  153. self.servers = [server]
  154. if sockets is None:
  155. self._log_started_message(listeners)
  156. else:
  157. # We're most likely running multiple workers, so a message has already been
  158. # logged by `config.bind_socket()`.
  159. pass # pragma: full coverage
  160. self.started = True
  161. def _log_started_message(self, listeners: Sequence[socket.SocketType]) -> None:
  162. config = self.config
  163. if config.fd is not None: # pragma: py-win32
  164. sock = listeners[0]
  165. logger.info(
  166. "Uvicorn running on socket %s (Press CTRL+C to quit)",
  167. sock.getsockname(),
  168. )
  169. elif config.uds is not None: # pragma: py-win32
  170. logger.info("Uvicorn running on unix socket %s (Press CTRL+C to quit)", config.uds)
  171. else:
  172. addr_format = "%s://%s:%d"
  173. host = "0.0.0.0" if config.host is None else config.host
  174. if ":" in host:
  175. # It's an IPv6 address.
  176. addr_format = "%s://[%s]:%d"
  177. port = config.port
  178. if port == 0:
  179. port = listeners[0].getsockname()[1]
  180. protocol_name = "https" if config.ssl else "http"
  181. message = f"Uvicorn running on {addr_format} (Press CTRL+C to quit)"
  182. color_message = "Uvicorn running on " + style(addr_format, bold=True) + " (Press CTRL+C to quit)"
  183. logger.info(
  184. message,
  185. protocol_name,
  186. host,
  187. port,
  188. extra={"color_message": color_message},
  189. )
  190. async def main_loop(self) -> None:
  191. counter = 0
  192. should_exit = await self.on_tick(counter)
  193. while not should_exit:
  194. counter += 1
  195. counter = counter % 864000
  196. await asyncio.sleep(0.1)
  197. should_exit = await self.on_tick(counter)
  198. async def on_tick(self, counter: int) -> bool:
  199. # Update the default headers, once per second.
  200. if counter % 10 == 0:
  201. current_time = time.time()
  202. current_date = formatdate(current_time, usegmt=True).encode()
  203. if self.config.date_header:
  204. date_header = [(b"date", current_date)]
  205. else:
  206. date_header = []
  207. self.server_state.default_headers = date_header + self.config.encoded_headers
  208. # Callback to `callback_notify` once every `timeout_notify` seconds.
  209. if self.config.callback_notify is not None:
  210. if current_time - self.last_notified > self.config.timeout_notify: # pragma: full coverage
  211. self.last_notified = current_time
  212. await self.config.callback_notify()
  213. # Determine if we should exit.
  214. if self.should_exit:
  215. return True
  216. max_requests = self.limit_max_requests
  217. if max_requests is not None and self.server_state.total_requests >= max_requests:
  218. logger.info("Maximum request limit of %d exceeded. Terminating process.", max_requests)
  219. return True
  220. return False
  221. async def shutdown(self, sockets: list[socket.socket] | None = None) -> None:
  222. logger.info("Shutting down")
  223. # Stop accepting new connections.
  224. for server in self.servers:
  225. server.close()
  226. for sock in sockets or []:
  227. sock.close() # pragma: full coverage
  228. # Request shutdown on all existing connections.
  229. for connection in list(self.server_state.connections):
  230. connection.shutdown()
  231. await asyncio.sleep(0.1)
  232. # When 3.10 is not supported anymore, use `async with asyncio.timeout(...):`.
  233. try:
  234. await asyncio.wait_for(
  235. self._wait_tasks_to_complete(),
  236. timeout=self.config.timeout_graceful_shutdown,
  237. )
  238. except asyncio.TimeoutError:
  239. logger.error(
  240. "Cancel %s running task(s), timeout graceful shutdown exceeded",
  241. len(self.server_state.tasks),
  242. )
  243. for t in self.server_state.tasks:
  244. t.cancel(msg="Task cancelled, timeout graceful shutdown exceeded")
  245. # Send the lifespan shutdown event, and wait for application shutdown.
  246. if not self.force_exit:
  247. await self.lifespan.shutdown()
  248. async def _wait_tasks_to_complete(self) -> None:
  249. # Wait for existing connections to finish sending responses.
  250. if self.server_state.connections and not self.force_exit:
  251. msg = "Waiting for connections to close. (CTRL+C to force quit)"
  252. logger.info(msg)
  253. while self.server_state.connections and not self.force_exit:
  254. await asyncio.sleep(0.1)
  255. # Wait for existing tasks to complete.
  256. if self.server_state.tasks and not self.force_exit:
  257. msg = "Waiting for background tasks to complete. (CTRL+C to force quit)"
  258. logger.info(msg)
  259. while self.server_state.tasks and not self.force_exit:
  260. await asyncio.sleep(0.1)
  261. for server in self.servers:
  262. await server.wait_closed()
  263. @contextlib.contextmanager
  264. def capture_signals(self) -> Generator[None, None, None]:
  265. # Signals can only be listened to from the main thread.
  266. if threading.current_thread() is not threading.main_thread():
  267. yield
  268. return
  269. # always use signal.signal, even if loop.add_signal_handler is available
  270. # this allows to restore previous signal handlers later on
  271. original_handlers = {sig: signal.signal(sig, self.handle_exit) for sig in HANDLED_SIGNALS}
  272. try:
  273. yield
  274. finally:
  275. for sig, handler in original_handlers.items():
  276. signal.signal(sig, handler)
  277. # If we did gracefully shut down due to a signal, try to
  278. # trigger the expected behaviour now; multiple signals would be
  279. # done LIFO, see https://stackoverflow.com/questions/48434964
  280. for captured_signal in reversed(self._captured_signals):
  281. signal.raise_signal(captured_signal)
  282. def handle_exit(self, sig: int, frame: FrameType | None) -> None:
  283. self._captured_signals.append(sig)
  284. if self.should_exit and sig == signal.SIGINT:
  285. self.force_exit = True # pragma: full coverage
  286. else:
  287. self.should_exit = True