v0.152.0
  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