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