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 os
 11import socket
 12import ssl
 13import stat
 14import sys
 15import time
 16from typing import TYPE_CHECKING
 17
 18from plain.logs import get_framework_logger
 19
 20from . import util
 21
 22if TYPE_CHECKING:
 23    from .app import ServerApplication
 24
 25log = get_framework_logger()
 26
 27# Maximum number of pending connections in the socket listen queue
 28BACKLOG = 2048
 29
 30
 31class BaseSocket:
 32    FAMILY: socket.AddressFamily
 33
 34    def __init__(
 35        self,
 36        address: tuple[str, int] | str,
 37        *,
 38        is_ssl: bool = False,
 39        fd: int | None = None,
 40    ) -> None:
 41        self.is_ssl = is_ssl
 42        self.cfg_addr = address
 43        if fd is None:
 44            sock = socket.socket(self.FAMILY, socket.SOCK_STREAM)
 45            bound = False
 46        else:
 47            sock = socket.fromfd(fd, self.FAMILY, socket.SOCK_STREAM)
 48            os.close(fd)
 49            bound = True
 50
 51        self.sock: socket.socket | None = self.set_options(sock, bound=bound)
 52
 53    def __str__(self) -> str:
 54        assert self.sock is not None, "Socket is closed"
 55        return f"<socket {self.sock.fileno()}>"
 56
 57    def __getattr__(self, name: str) -> object:
 58        return getattr(self.sock, name)
 59
 60    def accept(self) -> tuple[socket.socket, tuple[str, int] | str]:
 61        """Accept a connection. Returns (socket object, address)."""
 62        assert self.sock is not None, "Socket is closed"
 63        return self.sock.accept()
 64
 65    def fileno(self) -> int:
 66        """Return the socket's file descriptor."""
 67        assert self.sock is not None, "Socket is closed"
 68        return self.sock.fileno()
 69
 70    def setblocking(self, flag: bool) -> None:
 71        """Set blocking or non-blocking mode of the socket."""
 72        assert self.sock is not None, "Socket is closed"
 73        return self.sock.setblocking(flag)
 74
 75    def getsockname(self) -> tuple[str, int] | str:
 76        assert self.sock is not None, "Socket is closed"
 77        return self.sock.getsockname()
 78
 79    def set_options(self, sock: socket.socket, bound: bool = False) -> socket.socket:
 80        sock.setsockopt(socket.SOL_SOCKET, socket.SO_REUSEADDR, 1)
 81        if not bound:
 82            self.bind(sock)
 83        sock.setblocking(False)
 84        sock.listen(BACKLOG)
 85        return sock
 86
 87    def bind(self, sock: socket.socket) -> None:
 88        sock.bind(self.cfg_addr)
 89
 90    def close(self) -> None:
 91        if self.sock is None:
 92            return None
 93
 94        try:
 95            self.sock.close()
 96        except OSError as e:
 97            log.info("Error while closing socket", extra={"error": str(e)})
 98
 99        self.sock = None
100        return None
101
102
103class TCPSocket(BaseSocket):
104    FAMILY = socket.AF_INET
105
106    def __str__(self) -> str:
107        scheme = "https" if self.is_ssl else "http"
108
109        assert self.sock is not None, "Socket is closed"
110        addr = self.sock.getsockname()
111        return f"{scheme}://{addr[0]}:{addr[1]}"
112
113    def set_options(self, sock: socket.socket, bound: bool = False) -> socket.socket:
114        sock.setsockopt(socket.IPPROTO_TCP, socket.TCP_NODELAY, 1)
115        return super().set_options(sock, bound=bound)
116
117
118class TCP6Socket(TCPSocket):
119    FAMILY = socket.AF_INET6
120
121    def __str__(self) -> str:
122        assert self.sock is not None, "Socket is closed"
123        (host, port, _, _) = self.sock.getsockname()
124        return f"http://[{host}]:{port}"
125
126
127class UnixSocket(BaseSocket):
128    FAMILY = socket.AF_UNIX
129
130    def __init__(
131        self,
132        addr: str,
133        *,
134        is_ssl: bool = False,
135        fd: int | None = None,
136    ):
137        if fd is None:
138            try:
139                st = os.stat(addr)
140            except OSError as e:
141                if e.args[0] != errno.ENOENT:
142                    raise
143            else:
144                if stat.S_ISSOCK(st.st_mode):
145                    os.remove(addr)
146                else:
147                    raise ValueError(f"{addr!r} is not a socket")
148        super().__init__(addr, is_ssl=is_ssl, fd=fd)
149
150    def __str__(self) -> str:
151        return f"unix:{self.cfg_addr}"
152
153    def bind(self, sock: socket.socket) -> None:
154        sock.bind(self.cfg_addr)
155
156
157def _sock_type(addr: tuple[str, int] | str | bytes) -> type[BaseSocket]:
158    if isinstance(addr, tuple):
159        if util.is_ipv6(addr[0]):
160            sock_type = TCP6Socket
161        else:
162            sock_type = TCPSocket
163    elif isinstance(addr, str | bytes):
164        sock_type = UnixSocket
165    else:
166        raise TypeError(f"Unable to create socket from: {addr!r}")
167    return sock_type
168
169
170def create_sockets(app: ServerApplication) -> list[BaseSocket]:
171    """
172    Create a new socket for the configured addresses.
173
174    If a configured address is a tuple then a TCP socket is created.
175    If it is a string, a Unix socket is created. Otherwise, a TypeError is
176    raised.
177    """
178    listeners = []
179
180    # check ssl config early to raise the error on startup
181    # only the certfile is needed since it can contains the keyfile
182    if app.certfile and not os.path.exists(app.certfile):
183        raise ValueError(f'certfile "{app.certfile}" does not exist')
184
185    if app.keyfile and not os.path.exists(app.keyfile):
186        raise ValueError(f'keyfile "{app.keyfile}" does not exist')
187
188    for addr in app.address:
189        sock_type = _sock_type(addr)
190        sock = None
191        for i in range(5):
192            try:
193                sock = sock_type(addr, is_ssl=app.is_ssl)
194            except OSError as e:
195                if e.args[0] == errno.EADDRINUSE:
196                    log.error("Connection in use", extra={"addr": str(addr)})
197                if e.args[0] == errno.EADDRNOTAVAIL:
198                    log.error("Invalid address", extra={"addr": str(addr)})
199                log.error(
200                    "Connection failed",
201                    extra={"addr": str(addr), "error": str(e)},
202                )
203                if i < 5:
204                    log.debug("Retrying in 1 second.")
205                    time.sleep(1)
206            else:
207                break
208
209        if sock is None:
210            log.error("Can't connect", extra={"addr": str(addr)})
211            sys.exit(1)
212
213        listeners.append(sock)
214
215    return listeners
216
217
218def close_sockets(listeners: list[BaseSocket], unlink: bool = True) -> None:
219    for sock in listeners:
220        sock_name = sock.getsockname()
221        sock.close()
222        if unlink and _sock_type(sock_name) is UnixSocket:
223            assert isinstance(sock_name, str)
224            os.unlink(sock_name)
225
226
227def ssl_context(certfile: str, keyfile: str | None) -> ssl.SSLContext:
228    context = ssl.create_default_context(ssl.Purpose.CLIENT_AUTH)
229    context.load_cert_chain(certfile=certfile, keyfile=keyfile)
230    context.set_alpn_protocols(["h2", "http/1.1"])
231    return context