_sockets.py 13 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218219220221222223224225226227228229230231232233234235236237238239240241242243244245246247248249250251252253254255256257258259260261262263264265266267268269270271272273274275276277278279280281282283284285286287288289290291292293294295296297298299300301302303304305306307308309310311312313314315316317318319320321322323324325326327328329330331332333334335336337338339340341342343344345346347348349350351352353354355356357358359360361362363364365366367368369370371372373374375376377378379380381382383384385386387388389390391392393394395396397398399400401
  1. from __future__ import annotations
  2. import errno
  3. import socket
  4. from abc import abstractmethod
  5. from collections.abc import Callable, Collection, Mapping
  6. from contextlib import AsyncExitStack
  7. from io import IOBase
  8. from ipaddress import IPv4Address, IPv6Address
  9. from socket import AddressFamily
  10. from typing import TYPE_CHECKING, Any, TypeAlias, TypeVar
  11. from .._core._eventloop import get_async_backend
  12. from .._core._typedattr import (
  13. TypedAttributeProvider,
  14. TypedAttributeSet,
  15. typed_attribute,
  16. )
  17. from ._streams import ByteStream, Listener, UnreliableObjectStream
  18. if TYPE_CHECKING:
  19. from ._tasks import TaskGroup
  20. IPAddressType: TypeAlias = str | IPv4Address | IPv6Address
  21. IPSockAddrType: TypeAlias = tuple[str, int]
  22. SockAddrType: TypeAlias = IPSockAddrType | str
  23. UDPPacketType: TypeAlias = tuple[bytes, IPSockAddrType]
  24. UNIXDatagramPacketType: TypeAlias = tuple[bytes, str]
  25. T_Retval = TypeVar("T_Retval")
  26. def _validate_socket(
  27. sock_or_fd: socket.socket | int,
  28. sock_type: socket.SocketKind,
  29. addr_family: socket.AddressFamily = socket.AF_UNSPEC,
  30. *,
  31. require_connected: bool = False,
  32. require_bound: bool = False,
  33. ) -> socket.socket:
  34. if isinstance(sock_or_fd, int):
  35. try:
  36. sock = socket.socket(fileno=sock_or_fd)
  37. except OSError as exc:
  38. if exc.errno == errno.ENOTSOCK:
  39. raise ValueError(
  40. "the file descriptor does not refer to a socket"
  41. ) from exc
  42. elif require_connected:
  43. raise ValueError("the socket must be connected") from exc
  44. elif require_bound:
  45. raise ValueError("the socket must be bound to a local address") from exc
  46. else:
  47. raise
  48. elif isinstance(sock_or_fd, socket.socket):
  49. sock = sock_or_fd
  50. else:
  51. raise TypeError(
  52. f"expected an int or socket, got {type(sock_or_fd).__qualname__} instead"
  53. )
  54. try:
  55. if require_connected:
  56. try:
  57. sock.getpeername()
  58. except OSError as exc:
  59. raise ValueError("the socket must be connected") from exc
  60. if require_bound:
  61. try:
  62. if sock.family in (socket.AF_INET, socket.AF_INET6):
  63. bound_addr = sock.getsockname()[1]
  64. else:
  65. bound_addr = sock.getsockname()
  66. except OSError:
  67. bound_addr = None
  68. if not bound_addr:
  69. raise ValueError("the socket must be bound to a local address")
  70. if addr_family != socket.AF_UNSPEC and sock.family != addr_family:
  71. raise ValueError(
  72. f"address family mismatch: expected {addr_family.name}, got "
  73. f"{sock.family.name}"
  74. )
  75. if sock.type != sock_type:
  76. raise ValueError(
  77. f"socket type mismatch: expected {sock_type.name}, got {sock.type.name}"
  78. )
  79. except BaseException:
  80. # Avoid ResourceWarning from the locally constructed socket object
  81. if isinstance(sock_or_fd, int):
  82. sock.detach()
  83. raise
  84. sock.setblocking(False)
  85. return sock
  86. class SocketAttribute(TypedAttributeSet):
  87. """
  88. .. attribute:: family
  89. :type: socket.AddressFamily
  90. the address family of the underlying socket
  91. .. attribute:: local_address
  92. :type: tuple[str, int] | str
  93. the local address the underlying socket is connected to
  94. .. attribute:: local_port
  95. :type: int
  96. for IP based sockets, the local port the underlying socket is bound to
  97. .. attribute:: raw_socket
  98. :type: socket.socket
  99. the underlying stdlib socket object
  100. .. attribute:: remote_address
  101. :type: tuple[str, int] | str
  102. the remote address the underlying socket is connected to
  103. .. attribute:: remote_port
  104. :type: int
  105. for IP based sockets, the remote port the underlying socket is connected to
  106. """
  107. family: AddressFamily = typed_attribute()
  108. local_address: SockAddrType = typed_attribute()
  109. local_port: int = typed_attribute()
  110. raw_socket: socket.socket = typed_attribute()
  111. remote_address: SockAddrType = typed_attribute()
  112. remote_port: int = typed_attribute()
  113. class _SocketProvider(TypedAttributeProvider):
  114. @property
  115. def extra_attributes(self) -> Mapping[Any, Callable[[], Any]]:
  116. from .._core._sockets import convert_ipv6_sockaddr as convert
  117. attributes: dict[Any, Callable[[], Any]] = {
  118. SocketAttribute.family: lambda: self._raw_socket.family,
  119. SocketAttribute.local_address: lambda: convert(
  120. self._raw_socket.getsockname()
  121. ),
  122. SocketAttribute.raw_socket: lambda: self._raw_socket,
  123. }
  124. try:
  125. peername: tuple[str, int] | None = convert(self._raw_socket.getpeername())
  126. except OSError:
  127. peername = None
  128. # Provide the remote address for connected sockets
  129. if peername is not None:
  130. attributes[SocketAttribute.remote_address] = lambda: peername
  131. # Provide local and remote ports for IP based sockets
  132. if self._raw_socket.family in (AddressFamily.AF_INET, AddressFamily.AF_INET6):
  133. attributes[SocketAttribute.local_port] = lambda: (
  134. self._raw_socket.getsockname()[1]
  135. )
  136. if peername is not None:
  137. remote_port = peername[1]
  138. attributes[SocketAttribute.remote_port] = lambda: remote_port
  139. return attributes
  140. @property
  141. @abstractmethod
  142. def _raw_socket(self) -> socket.socket:
  143. pass
  144. class SocketStream(ByteStream, _SocketProvider):
  145. """
  146. Transports bytes over a socket.
  147. Supports all relevant extra attributes from :class:`~SocketAttribute`.
  148. """
  149. @classmethod
  150. async def from_socket(cls, sock_or_fd: socket.socket | int) -> SocketStream:
  151. """
  152. Wrap an existing socket object or file descriptor as a socket stream.
  153. The newly created socket wrapper takes ownership of the socket being passed in.
  154. The existing socket must already be connected.
  155. :param sock_or_fd: a socket object or file descriptor
  156. :return: a socket stream
  157. """
  158. sock = _validate_socket(sock_or_fd, socket.SOCK_STREAM, require_connected=True)
  159. return await get_async_backend().wrap_stream_socket(sock)
  160. class UNIXSocketStream(SocketStream):
  161. @classmethod
  162. async def from_socket(cls, sock_or_fd: socket.socket | int) -> UNIXSocketStream:
  163. """
  164. Wrap an existing socket object or file descriptor as a UNIX socket stream.
  165. The newly created socket wrapper takes ownership of the socket being passed in.
  166. The existing socket must already be connected.
  167. :param sock_or_fd: a socket object or file descriptor
  168. :return: a UNIX socket stream
  169. """
  170. sock = _validate_socket(
  171. sock_or_fd, socket.SOCK_STREAM, socket.AF_UNIX, require_connected=True
  172. )
  173. return await get_async_backend().wrap_unix_stream_socket(sock)
  174. @abstractmethod
  175. async def send_fds(self, message: bytes, fds: Collection[int | IOBase]) -> None:
  176. """
  177. Send file descriptors along with a message to the peer.
  178. :param message: a non-empty bytestring
  179. :param fds: a collection of files (either numeric file descriptors or open file
  180. or socket objects)
  181. """
  182. @abstractmethod
  183. async def receive_fds(self, msglen: int, maxfds: int) -> tuple[bytes, list[int]]:
  184. """
  185. Receive file descriptors along with a message from the peer.
  186. :param msglen: length of the message to expect from the peer
  187. :param maxfds: maximum number of file descriptors to expect from the peer
  188. :return: a tuple of (message, file descriptors)
  189. """
  190. class SocketListener(Listener[SocketStream], _SocketProvider):
  191. """
  192. Listens to incoming socket connections.
  193. Supports all relevant extra attributes from :class:`~SocketAttribute`.
  194. """
  195. @classmethod
  196. async def from_socket(
  197. cls,
  198. sock_or_fd: socket.socket | int,
  199. ) -> SocketListener:
  200. """
  201. Wrap an existing socket object or file descriptor as a socket listener.
  202. The newly created listener takes ownership of the socket being passed in.
  203. :param sock_or_fd: a socket object or file descriptor
  204. :return: a socket listener
  205. """
  206. sock = _validate_socket(sock_or_fd, socket.SOCK_STREAM, require_bound=True)
  207. return await get_async_backend().wrap_listener_socket(sock)
  208. @abstractmethod
  209. async def accept(self) -> SocketStream:
  210. """Accept an incoming connection."""
  211. async def serve(
  212. self,
  213. handler: Callable[[SocketStream], Any],
  214. task_group: TaskGroup | None = None,
  215. ) -> None:
  216. from .. import create_task_group
  217. async with AsyncExitStack() as stack:
  218. if task_group is None:
  219. task_group = await stack.enter_async_context(create_task_group())
  220. while True:
  221. stream = await self.accept()
  222. task_group.start_soon(handler, stream)
  223. class UDPSocket(UnreliableObjectStream[UDPPacketType], _SocketProvider):
  224. """
  225. Represents an unconnected UDP socket.
  226. Supports all relevant extra attributes from :class:`~SocketAttribute`.
  227. """
  228. @classmethod
  229. async def from_socket(cls, sock_or_fd: socket.socket | int) -> UDPSocket:
  230. """
  231. Wrap an existing socket object or file descriptor as a UDP socket.
  232. The newly created socket wrapper takes ownership of the socket being passed in.
  233. The existing socket must be bound to a local address.
  234. :param sock_or_fd: a socket object or file descriptor
  235. :return: a UDP socket
  236. """
  237. sock = _validate_socket(sock_or_fd, socket.SOCK_DGRAM, require_bound=True)
  238. return await get_async_backend().wrap_udp_socket(sock)
  239. async def sendto(self, data: bytes, host: str, port: int) -> None:
  240. """
  241. Alias for :meth:`~.UnreliableObjectSendStream.send` ((data, (host, port))).
  242. """
  243. return await self.send((data, (host, port)))
  244. class ConnectedUDPSocket(UnreliableObjectStream[bytes], _SocketProvider):
  245. """
  246. Represents an connected UDP socket.
  247. Supports all relevant extra attributes from :class:`~SocketAttribute`.
  248. """
  249. @classmethod
  250. async def from_socket(cls, sock_or_fd: socket.socket | int) -> ConnectedUDPSocket:
  251. """
  252. Wrap an existing socket object or file descriptor as a connected UDP socket.
  253. The newly created socket wrapper takes ownership of the socket being passed in.
  254. The existing socket must already be connected.
  255. :param sock_or_fd: a socket object or file descriptor
  256. :return: a connected UDP socket
  257. """
  258. sock = _validate_socket(
  259. sock_or_fd,
  260. socket.SOCK_DGRAM,
  261. require_connected=True,
  262. )
  263. return await get_async_backend().wrap_connected_udp_socket(sock)
  264. class UNIXDatagramSocket(
  265. UnreliableObjectStream[UNIXDatagramPacketType], _SocketProvider
  266. ):
  267. """
  268. Represents an unconnected Unix datagram socket.
  269. Supports all relevant extra attributes from :class:`~SocketAttribute`.
  270. """
  271. @classmethod
  272. async def from_socket(
  273. cls,
  274. sock_or_fd: socket.socket | int,
  275. ) -> UNIXDatagramSocket:
  276. """
  277. Wrap an existing socket object or file descriptor as a UNIX datagram
  278. socket.
  279. The newly created socket wrapper takes ownership of the socket being passed in.
  280. :param sock_or_fd: a socket object or file descriptor
  281. :return: a UNIX datagram socket
  282. """
  283. sock = _validate_socket(sock_or_fd, socket.SOCK_DGRAM, socket.AF_UNIX)
  284. return await get_async_backend().wrap_unix_datagram_socket(sock)
  285. async def sendto(self, data: bytes, path: str) -> None:
  286. """Alias for :meth:`~.UnreliableObjectSendStream.send` ((data, path))."""
  287. return await self.send((data, path))
  288. class ConnectedUNIXDatagramSocket(UnreliableObjectStream[bytes], _SocketProvider):
  289. """
  290. Represents a connected Unix datagram socket.
  291. Supports all relevant extra attributes from :class:`~SocketAttribute`.
  292. """
  293. @classmethod
  294. async def from_socket(
  295. cls,
  296. sock_or_fd: socket.socket | int,
  297. ) -> ConnectedUNIXDatagramSocket:
  298. """
  299. Wrap an existing socket object or file descriptor as a connected UNIX datagram
  300. socket.
  301. The newly created socket wrapper takes ownership of the socket being passed in.
  302. The existing socket must already be connected.
  303. :param sock_or_fd: a socket object or file descriptor
  304. :return: a connected UNIX datagram socket
  305. """
  306. sock = _validate_socket(
  307. sock_or_fd, socket.SOCK_DGRAM, socket.AF_UNIX, require_connected=True
  308. )
  309. return await get_async_backend().wrap_connected_unix_datagram_socket(sock)