1from __future__ import annotations
2
3#
4#
5# This file is part of gunicorn released under the MIT license.
6# See the LICENSE for more information.
7#
8# Vendored and modified for Plain.
9import errno
10import multiprocessing
11import os
12import signal
13import sys
14import threading
15import time
16from dataclasses import dataclass, field
17from typing import TYPE_CHECKING
18
19import plain.runtime
20from plain.logs import get_framework_logger
21from plain.runtime import settings
22
23from . import sock
24from .errors import APP_LOAD_ERROR, WORKER_BOOT_ERROR, HaltServer
25from .workers.entry import worker_main
26from .workers.worker import check_worker_config
27from .workers.workertmp import WorkerHeartbeat
28
29if TYPE_CHECKING:
30 from .app import ServerApplication
31
32
33@dataclass
34class WorkerInfo:
35 process: multiprocessing.process.BaseProcess
36 heartbeat: WorkerHeartbeat
37 age: int
38 spawned_at: float = field(default_factory=time.monotonic)
39 aborted: bool = field(default=False)
40
41
42class Arbiter:
43 """
44 Arbiter maintains the worker processes alive. It launches or
45 kills them if needed.
46 """
47
48 def __init__(self, app: ServerApplication):
49 os.environ["SERVER_SOFTWARE"] = f"plain/{plain.runtime.__version__}"
50
51 self.app = app
52 self.log = get_framework_logger()
53 self.num_workers: int = app.workers
54 self.timeout: int = app.timeout
55 self.pid: int = os.getpid()
56 self.worker_age: int = 0
57 self._workers: dict[int, WorkerInfo] = {}
58 self._listeners: list[sock.BaseSocket] = []
59 self._shutdown_event = threading.Event()
60 self._graceful_shutdown = True
61 self._halt_error: HaltServer | None = None
62 self._last_logged_active_worker_count: int | None = None
63 self._mp_context = multiprocessing.get_context("spawn")
64
65 def run(self) -> None:
66 """Main supervisor loop."""
67 self._start()
68
69 try:
70 self.manage_workers()
71
72 while not self._shutdown_event.is_set():
73 self.reap_workers()
74 if self._halt_error:
75 raise self._halt_error
76 self.murder_workers()
77 self.manage_workers()
78 self._shutdown_event.wait(timeout=1.0)
79
80 self._halt(graceful=self._graceful_shutdown)
81 except KeyboardInterrupt:
82 self._halt(graceful=False)
83 except HaltServer as inst:
84 self._halt(reason=inst.reason, exit_status=inst.exit_status)
85 except SystemExit:
86 raise
87 except Exception:
88 self.log.error("Unhandled exception in main loop", exc_info=True)
89 self._stop(graceful=False)
90 sys.exit(-1)
91
92 def _start(self) -> None:
93 """Initialize the arbiter. Start listening."""
94 # SIGTERM = graceful shutdown, SIGINT/SIGQUIT = immediate shutdown
95 signal.signal(signal.SIGTERM, self._handle_signal)
96 signal.signal(signal.SIGINT, self._handle_hard_stop)
97 signal.signal(signal.SIGQUIT, self._handle_hard_stop)
98 signal.signal(signal.SIGUSR1, self._handle_memory_signal)
99
100 self._listeners = sock.create_sockets(self.app)
101
102 listeners_str = ",".join([str(lnr) for lnr in self._listeners])
103 self.log.info(
104 "Plain server started",
105 extra={
106 "address": listeners_str,
107 "pid": self.pid,
108 "workers": self.num_workers,
109 "threads": self.app.threads,
110 "version": plain.runtime.__version__,
111 },
112 )
113
114 from plain.runtime import settings
115
116 check_worker_config(self.app.threads, settings.SERVER_CONNECTIONS, self.log)
117
118 def _handle_memory_signal(self, sig: int, frame: object) -> None:
119 """Forward SIGUSR1 to all workers for memory profiling."""
120 self._kill_workers(signal.SIGUSR1)
121
122 def _handle_signal(self, sig: int, frame: object) -> None:
123 self._shutdown_event.set()
124
125 def _handle_hard_stop(self, sig: int, frame: object) -> None:
126 self._graceful_shutdown = False
127 self._shutdown_event.set()
128
129 def _halt(
130 self, reason: str | None = None, exit_status: int = 0, graceful: bool = True
131 ) -> None:
132 """Halt arbiter."""
133 self._stop(graceful=graceful)
134
135 log_func = self.log.info if exit_status == 0 else self.log.error
136 log_func("Shutting down: Master")
137 if reason is not None:
138 log_func("Shutting down", extra={"reason": reason})
139
140 sys.exit(exit_status)
141
142 def _stop(self, graceful: bool = True) -> None:
143 """Stop workers."""
144 sock.close_sockets(self._listeners, unlink=True)
145 self._listeners = []
146
147 sig = signal.SIGTERM if graceful else signal.SIGQUIT
148 # Monotonic, to share one clock with the workers' drain deadlines
149 # (WorkerHeartbeat relies on CLOCK_MONOTONIC being system-wide).
150 limit = time.monotonic() + settings.SERVER_GRACEFUL_TIMEOUT
151
152 # This shutdown ends in SIGKILL at the limit — publish it so the
153 # workers cap their drains and finish teardown first. (Retirement
154 # SIGTERMs in manage_workers have no SIGKILL follower and don't
155 # set this.)
156 for info in self._workers.values():
157 info.heartbeat.set_kill_deadline(limit)
158
159 # Instruct the workers to exit
160 self._kill_workers(sig)
161
162 # Wait until the graceful timeout
163 while self._workers and time.monotonic() < limit:
164 self.reap_workers()
165 time.sleep(0.1)
166
167 self._kill_workers(signal.SIGKILL)
168
169 # Join and close all remaining processes
170 for pid in list(self._workers):
171 info = self._workers.pop(pid)
172 info.process.join(timeout=5)
173 info.heartbeat.close()
174 info.process.close()
175
176 def murder_workers(self) -> None:
177 """Kill workers that have stopped heartbeating."""
178 if not self.timeout:
179 return
180
181 now = time.monotonic()
182 for pid, info in list(self._workers.items()):
183 # Don't kill workers that haven't had enough time to boot
184 # and start heartbeating (spawn is slower than fork).
185 if now - info.spawned_at < self.timeout:
186 continue
187
188 try:
189 if now - info.heartbeat.last_update() <= self.timeout:
190 continue
191 except (OSError, ValueError):
192 continue
193
194 if not info.aborted:
195 self.log.critical("WORKER TIMEOUT", extra={"pid": pid})
196 info.aborted = True
197 self._kill_worker(pid, signal.SIGABRT)
198 else:
199 self._kill_worker(pid, signal.SIGKILL)
200
201 def reap_workers(self) -> None:
202 """
203 Reap dead workers and log exit reasons.
204 Sets self._halt_error if a worker failed to boot.
205 """
206 for pid in list(self._workers):
207 info = self._workers[pid]
208 if info.process.is_alive():
209 continue
210
211 exitcode = info.process.exitcode
212 if exitcode is None:
213 continue
214
215 if exitcode > 0:
216 self.log.error(
217 "Worker exited with error code",
218 extra={"pid": pid, "exitcode": exitcode},
219 )
220
221 if exitcode == WORKER_BOOT_ERROR and self._halt_error is None:
222 self._halt_error = HaltServer(
223 "Worker failed to boot.", WORKER_BOOT_ERROR
224 )
225 elif exitcode == APP_LOAD_ERROR and self._halt_error is None:
226 self._halt_error = HaltServer("App failed to load.", APP_LOAD_ERROR)
227 elif exitcode < 0:
228 # Negative exit codes mean the worker was killed by a signal
229 try:
230 sig_name = signal.Signals(-exitcode).name
231 except ValueError:
232 sig_name = f"signal {-exitcode}"
233 note = "Perhaps out of memory?" if -exitcode == signal.SIGKILL else ""
234 ctx = {"pid": pid, "signal": sig_name}
235 if note:
236 ctx["note"] = note
237 if -exitcode == signal.SIGTERM:
238 self.log.info("Worker was sent signal", extra=ctx)
239 else:
240 self.log.error("Worker was sent signal", extra=ctx)
241
242 info.heartbeat.close()
243 info.process.join(timeout=0)
244 info.process.close()
245 del self._workers[pid]
246
247 def manage_workers(self) -> None:
248 """Maintain the number of workers by spawning or killing as required.
249
250 Retiring workers (hit max_requests) are not counted toward the target.
251 Replacements are pre-spawned so the retiring worker can keep serving
252 traffic until the replacement is ready, avoiding dropped connections.
253 """
254 active_pids = []
255 retiring_pids = []
256
257 for pid, info in self._workers.items():
258 if info.heartbeat.is_retiring():
259 retiring_pids.append(pid)
260 else:
261 active_pids.append(pid)
262
263 # Spawn replacements for retiring workers (and any other shortfall)
264 active_count = len(active_pids)
265 while active_count < self.num_workers:
266 self._spawn_worker()
267 active_count += 1
268
269 # Kill excess non-retiring workers
270 if active_count > self.num_workers:
271 workers = sorted(
272 [(pid, self._workers[pid]) for pid in active_pids],
273 key=lambda w: w[1].age,
274 )
275 while len(workers) > self.num_workers:
276 (pid, _) = workers.pop(0)
277 self._kill_worker(pid, signal.SIGTERM)
278
279 # Once enough non-retiring workers are ready (heartbeating),
280 # tell retiring workers to shut down gracefully.
281 ready_count = sum(
282 1
283 for pid, info in self._workers.items()
284 if not info.heartbeat.is_retiring()
285 and info.heartbeat.last_update() > info.spawned_at
286 )
287 if ready_count >= self.num_workers:
288 for pid in retiring_pids:
289 if pid in self._workers:
290 self.log.info(
291 "Replacement ready, shutting down retiring worker",
292 extra={"pid": pid},
293 )
294 self._kill_worker(pid, signal.SIGTERM)
295
296 active_worker_count = len(active_pids)
297 if self._last_logged_active_worker_count != active_worker_count:
298 self._last_logged_active_worker_count = active_worker_count
299 self.log.debug(
300 f"{active_worker_count} workers",
301 extra={
302 "metric": "plain.server.workers",
303 "value": active_worker_count,
304 "mtype": "gauge",
305 },
306 )
307
308 def _spawn_worker(self) -> None:
309 self.worker_age += 1
310 heartbeat = WorkerHeartbeat(self._mp_context)
311
312 # Serialize listener info for the spawned process.
313 # Raw socket objects are pickled via multiprocessing (SCM_RIGHTS on Unix).
314 listener_data = [
315 (listener.sock, listener.cfg_addr, listener.FAMILY, listener.is_ssl)
316 for listener in self._listeners
317 ]
318
319 process = self._mp_context.Process(
320 target=worker_main,
321 args=(
322 self.worker_age,
323 listener_data,
324 self.app,
325 self.timeout / 2.0,
326 heartbeat,
327 ),
328 )
329 process.start()
330 assert process.pid is not None
331 self._workers[process.pid] = WorkerInfo(process, heartbeat, self.worker_age)
332
333 def _kill_workers(self, sig: int) -> None:
334 """Kill all workers with the signal `sig`."""
335 for pid in list(self._workers.keys()):
336 self._kill_worker(pid, sig)
337
338 def _kill_worker(self, pid: int, sig: int) -> None:
339 """Kill a worker."""
340 try:
341 os.kill(pid, sig)
342 except OSError as e:
343 if e.errno == errno.ESRCH:
344 try:
345 info = self._workers.pop(pid)
346 info.heartbeat.close()
347 info.process.close()
348 except (KeyError, OSError):
349 pass
350 return
351 raise