diff --git a/HISTORY.rst b/HISTORY.rst index a7100f5..33540e3 100644 --- a/HISTORY.rst +++ b/HISTORY.rst @@ -10,7 +10,8 @@ aiofastnet Release History * Harden logic against exceptions in SSLTransport constructor * Added create_datagram_endpoint -* Fixed Protocol.connection_lost may not be called if failure and abort happenes during start_tls +* Added connect_accepted_socket +* Fixed Protocol.connection_lost may not be called if failure and abort happens during start_tls 0.21.0 ------------------ diff --git a/README.md b/README.md index d0730cd..6fd4711 100644 --- a/README.md +++ b/README.md @@ -69,6 +69,7 @@ Source: [examples/benchmark_threaded.py](https://github.com/tarasko/aiofastnet/b `aiofastnet` provides drop-in, highly efficient C/Cython replacements for asyncio's: - [`loop.create_connection()`](https://docs.python.org/3/library/asyncio-eventloop.html#asyncio.loop.create_connection) +- [`loop.connect_accepted_socket()`](https://docs.python.org/3/library/asyncio-eventloop.html#asyncio.loop.connect_accepted_socket) - [`loop.open_connection()`](https://docs.python.org/3/library/asyncio-stream.html#asyncio.open_connection) - [`loop.create_unix_connection()`](https://docs.python.org/3/library/asyncio-eventloop.html#asyncio.loop.create_unix_connection) - [`loop.open_unix_connection()`](https://docs.python.org/3/library/asyncio-stream.html#asyncio.open_unix_connection) diff --git a/aiofastnet/__init__.py b/aiofastnet/__init__.py index e0a48bb..38f87c6 100644 --- a/aiofastnet/__init__.py +++ b/aiofastnet/__init__.py @@ -1,5 +1,6 @@ import socket +from .api_connect_accepted_socket import connect_accepted_socket from .api_create_connection import create_connection from .api_create_datagram_endpoint import create_datagram_endpoint from .api_create_server import create_server @@ -20,6 +21,7 @@ 'Protocol', 'Transport', 'aiofn_is_buffered_protocol', + 'connect_accepted_socket', 'create_connection', 'create_datagram_endpoint', 'create_server', diff --git a/aiofastnet/__init__.pyi b/aiofastnet/__init__.pyi index e8e5c29..b507081 100644 --- a/aiofastnet/__init__.pyi +++ b/aiofastnet/__init__.pyi @@ -63,6 +63,18 @@ async def create_connection( all_errors: bool = ..., ) -> tuple[asyncio.Transport, _ProtocolT]: ... +async def connect_accepted_socket( + loop: asyncio.AbstractEventLoop, + protocol_factory: Callable[[], _ProtocolT], + sock: socket.socket, + *, + ssl: bool | ssl.SSLContext | None = ..., + ssl_handshake_timeout: float | None = ..., + ssl_shutdown_timeout: float | None = ..., + ssl_incoming_bio_size: int | None = ..., + ssl_outgoing_bio_size: int | None = ..., +) -> tuple[asyncio.Transport, _ProtocolT]: ... + async def create_datagram_endpoint( loop: asyncio.AbstractEventLoop, protocol_factory: Callable[[], _DatagramProtocolT], diff --git a/aiofastnet/api_connect_accepted_socket.py b/aiofastnet/api_connect_accepted_socket.py new file mode 100644 index 0000000..904d13c --- /dev/null +++ b/aiofastnet/api_connect_accepted_socket.py @@ -0,0 +1,46 @@ +# Portions of this file are derived from CPython's asyncio sources +# (notably asyncio.base_events and asyncio.selector_events). +# Copyright (c) Python Software Foundation. +# Licensed under the Python Software Foundation License Version 2. +# See LICENSES/PSF-2.0.txt and THIRD_PARTY_NOTICES for details. + +import socket + +from .api_utils import _check_ssl_socket, _create_connection_transport, _logger, _validate_bio_size, _validate_ssl_timeout + + +async def connect_accepted_socket( + loop, + protocol_factory, + sock, + *, + ssl=None, + ssl_handshake_timeout=None, + ssl_shutdown_timeout=None, + ssl_incoming_bio_size=None, + ssl_outgoing_bio_size=None, +): + if sock.type != socket.SOCK_STREAM: + raise ValueError(f"A Stream Socket was expected, got {sock!r}") + + ssl_handshake_timeout = _validate_ssl_timeout("ssl_handshake_timeout", ssl_handshake_timeout, ssl) + ssl_shutdown_timeout = _validate_ssl_timeout("ssl_shutdown_timeout", ssl_shutdown_timeout, ssl) + ssl_incoming_bio_size = _validate_bio_size("ssl_incoming_bio_size", ssl_incoming_bio_size, ssl) + ssl_outgoing_bio_size = _validate_bio_size("ssl_outgoing_bio_size", ssl_outgoing_bio_size, ssl) + + _check_ssl_socket(sock) + + transport, protocol = await _create_connection_transport( + loop, sock, protocol_factory, ssl, "", + server_side=True, + ssl_handshake_timeout=ssl_handshake_timeout, + ssl_shutdown_timeout=ssl_shutdown_timeout, + ssl_incoming_bio_size=ssl_incoming_bio_size, + ssl_outgoing_bio_size=ssl_outgoing_bio_size, + ) + if loop.get_debug(): + # Get the socket from the transport because SSL transport closes + # the old socket and creates a new SSL socket + sock = transport.get_extra_info("socket") + _logger.debug("%r handled: (%r, %r)", sock, transport, protocol) + return transport, protocol diff --git a/tests/test_smoke.py b/tests/test_smoke.py index 1e3a21d..c353bc5 100644 --- a/tests/test_smoke.py +++ b/tests/test_smoke.py @@ -13,7 +13,33 @@ import aiofastnet from aiofastnet.transport import Protocol, SocketTransport, Transport -from tests.utils import UDP_MAX_PAYLOAD_SIZE, AsyncClient, SomeException, TestClient, TestServer, _logger, make_test_ssl_contexts, sendfile, start_tls +from tests.utils import ( + UDP_MAX_PAYLOAD_SIZE, + AsyncClient, + EchoServerProtocol, + SocketPair, + SomeException, + TestClient, + TestServer, + _logger, + make_test_ssl_contexts, + sendfile, + start_tls, +) + + +async def test_echo_socketpair(conn_type_plus_udp): + msg_size = 1024 + payload = b"x" * msg_size + + async with SocketPair(ct=conn_type_plus_udp, + server_protocol_factory=EchoServerProtocol, + client_server_hostname="127.0.0.1") as (_server, client): + client.write(payload) + echoed = await client.readn(msg_size) + assert echoed == payload + client.close() + await client.wait_closed() @pytest.mark.parametrize("msg_size", [1, 2, 3, 4, 5, 6, 7, 8, 29, 64, 256 * 1024, 6 * 1024 * 1024]) @@ -35,6 +61,18 @@ async def test_echo(all_loops, msg_size, conn_type_plus_udp, buffered_protocol): await client.wait_closed() +@pytest.mark.parametrize("msg_size", [1, 32, 64, 256 * 1024, 6 * 1024 * 1024, 20 * 1024 * 1024]) +@pytest.mark.parametrize("num_lines", [1, 32, 4000]) +async def test_echo_writelines(all_loops, msg_size, num_lines, conn_type, buffered_protocol): + payload = b"x" * msg_size + + async with TestServer(ct=conn_type, is_buffered=buffered_protocol) as server: + async with TestClient(server, ct=conn_type, is_buffered=buffered_protocol) as client: + client.write_in_lines(payload, num_lines) + echoed = await client.readn(msg_size, 4.0) + assert echoed == payload + + async def test_ktls_enabled(ktls_conn_type): async with TestServer(ct=ktls_conn_type) as server: async with TestClient(server, ct=ktls_conn_type) as client: @@ -78,18 +116,6 @@ async def test_ssl_membio_enabled(selector_loop, ssl_conn_type): assert server_client.transport.get_extra_info("ssl_outgoing_use_membio") is expected -@pytest.mark.parametrize("msg_size", [1, 32, 64, 256 * 1024, 6 * 1024 * 1024, 20 * 1024 * 1024]) -@pytest.mark.parametrize("num_lines", [1, 32, 4000]) -async def test_echo_writelines(all_loops, msg_size, num_lines, conn_type, buffered_protocol): - payload = b"x" * msg_size - - async with TestServer(ct=conn_type, is_buffered=buffered_protocol) as server: - async with TestClient(server, ct=conn_type, is_buffered=buffered_protocol) as client: - client.write_in_lines(payload, num_lines) - echoed = await client.readn(msg_size, 4.0) - assert echoed == payload - - async def test_write_huge_close(all_loops, conn_type): if os.name == 'nt' and isinstance(asyncio.get_running_loop(), asyncio.ProactorEventLoop) and sys.version_info < (3, 11): pytest.skip("ProactorEventLoop in 3.9 and 3.10 had issues with connection closing") diff --git a/tests/utils.py b/tests/utils.py index e95d41f..8fdaf21 100644 --- a/tests/utils.py +++ b/tests/utils.py @@ -7,7 +7,7 @@ import sys import tempfile import weakref -from contextlib import ExitStack, asynccontextmanager, contextmanager +from contextlib import AsyncExitStack, ExitStack, asynccontextmanager, contextmanager from dataclasses import dataclass from logging import getLogger from pathlib import Path @@ -84,12 +84,34 @@ async def create_datagram_endpoint(loop, *args, **kwargs): return await aiofastnet.create_datagram_endpoint(loop, *args, **kwargs) +async def connect_accepted_socket(loop, *args, **kwargs): + if NO_AIOFN: + return await loop.connect_accepted_socket(*args, **kwargs) + else: + return await aiofastnet.connect_accepted_socket(loop, *args, **kwargs) + + class EchoServerProtocol(asyncio.Protocol, asyncio.BufferedProtocol): - def __init__(self, clients: set, client_waiters: list[Any], is_buffered: bool): + transport: asyncio.Transport | None + _is_buffered: bool + _clients: set | None + _client_waiters: list[asyncio.Future] | None + _ssl_layer_num: int + _read_buffer: bytearray + + def __init__(self, + is_buffered: bool=False, + clients: set | None=None, + client_waiters: list[asyncio.Future] | None=None): self.transport = None + self._is_buffered = is_buffered self._clients = clients self._client_waiters = client_waiters - self._is_buffered = is_buffered + if clients is not None: + assert self._client_waiters is not None + else: + assert self._client_waiters is None + self._ssl_layer_num = 0 self._read_buffer = bytearray(b"X") * (128*1024) def is_buffered_protocol(self): @@ -97,19 +119,24 @@ def is_buffered_protocol(self): def connection_made(self, transport): _logger.debug("EchoServer.connection_made") - self._clients.add(weakref.ref(self)) self.transport = transport + ssl_protocol = self.transport.get_extra_info('ssl_protocol') if ssl_protocol is not None and hasattr(ssl_protocol, '_allow_renegotiation'): ssl_protocol._allow_renegotiation() - for w in self._client_waiters: - if not w.done(): - w.set_result(None) - self._client_waiters.clear() + + if self._clients is not None: + self._clients.add(weakref.ref(self)) + assert self._client_waiters is not None + for w in self._client_waiters: + if not w.done(): + w.set_result(None) + self._client_waiters.clear() def connection_lost(self, exc): _logger.debug("EchoServer.connection_lost, exc=%s", exc) - self._clients.remove(weakref.ref(self)) + if self._clients is not None: + self._clients.remove(weakref.ref(self)) def get_buffer(self, hint): return memoryview(self._read_buffer) @@ -135,6 +162,20 @@ def resume_writing(self): def eof_received(self): _logger.debug("EchoServer.eof_received") + async def start_tls(self, ssl_context, + ssl_handshake_timeout=None, ssl_shutdown_timeout=None): + self.transport = await start_tls( + asyncio.get_running_loop(), + self.transport, + self, + ssl_context, + server_side=True, + ssl_handshake_timeout=ssl_handshake_timeout, + ssl_shutdown_timeout=ssl_shutdown_timeout, + ) + _logger.debug("Server start_tls #%d completed", self._ssl_layer_num) + self._ssl_layer_num += 1 + class AsyncClient(asyncio.Protocol, asyncio.BufferedProtocol): transport: asyncio.Transport | None @@ -243,7 +284,8 @@ def connection_lost(self, exc): else: self._closed_fut.set_result(None) if self._readn_waiter is not None: - self._readn_waiter[1].set_exception(ConnectionResetError()) + if not self._readn_waiter[1].done(): + self._readn_waiter[1].set_exception(ConnectionResetError()) self._readn_waiter = None if self._write_resumed_fut is not None: self._write_resumed_fut.set_exception(ConnectionResetError()) @@ -566,7 +608,7 @@ async def TestServer(protocol_factory=None, client_waiters = [] if protocol_factory is None: def protocol_factory(): - return EchoServerProtocol(clients, client_waiters, is_buffered) + return EchoServerProtocol(is_buffered, clients, client_waiters) with ExitStack() as stack: if ct.name == "udp": @@ -639,7 +681,8 @@ async def TestClient(server_or_host=None, port=None, protocol_factory=AsyncClient, ssl_handshake_timeout=None, ssl_shutdown_timeout=None, - sock=None): + sock=None, + sock_server_side=False): if ct is None: ct = ConnectionType("tcp") if sock is not None: @@ -665,15 +708,10 @@ def client_protocol_factory(): return protocol try: - if ct.name == "unix": - transport, client = await create_unix_connection( - loop, - client_protocol_factory, - path=path if path is not None else host, - ) - elif ct.name == "udp": + if ct.name == "udp": if is_buffered: pytest.skip("UDP protocol is always simple") + if sock is not None: transport, client = await create_datagram_endpoint( loop, @@ -686,31 +724,57 @@ def client_protocol_factory(): client_protocol_factory, remote_addr=(host, port), ) - elif ct.use_start_tls or ct.client_ssl_context is None: - transport, client = await create_connection( - loop, - client_protocol_factory, - host=host, - port=port, - ) else: - transport, client = await create_connection( - loop, - client_protocol_factory, - host=host, - port=port, - ssl=ct.client_ssl_context, - server_hostname=server_hostname, - ssl_handshake_timeout=ssl_handshake_timeout, - ssl_shutdown_timeout=ssl_shutdown_timeout - ) - if ct.use_start_tls: - await client.start_tls(ct.client_ssl_context, - server_hostname=server_hostname, - ssl_handshake_timeout=ssl_handshake_timeout, - ssl_shutdown_timeout=ssl_shutdown_timeout - ) - + if not sock_server_side or sock is None: + if ct.name == "unix": + transport, client = await create_unix_connection( + loop, + client_protocol_factory, + path=path if path is not None else host, + sock=sock, + ) + elif ct.use_start_tls or ct.client_ssl_context is None: + transport, client = await create_connection( + loop, + client_protocol_factory, + host=host, + port=port, + sock=sock + ) + else: + transport, client = await create_connection( + loop, + client_protocol_factory, + host=host, + port=port, + sock=sock, + ssl=ct.client_ssl_context, + server_hostname=server_hostname, + ssl_handshake_timeout=ssl_handshake_timeout, + ssl_shutdown_timeout=ssl_shutdown_timeout + ) + if ct.use_start_tls: + await client.start_tls(ct.client_ssl_context, + server_hostname=server_hostname, + ssl_handshake_timeout=ssl_handshake_timeout, + ssl_shutdown_timeout=ssl_shutdown_timeout + ) + else: + if ct.use_start_tls or ct.client_ssl_context is None: + transport, client = await connect_accepted_socket(loop, client_protocol_factory, sock) + else: + transport, client = await connect_accepted_socket( + loop, + client_protocol_factory, + sock, + ssl=ct.server_ssl_context, + ssl_handshake_timeout=ssl_handshake_timeout, + ssl_shutdown_timeout=ssl_shutdown_timeout, + ) + if ct.use_start_tls: + await client.start_tls(ct.server_ssl_context, + ssl_handshake_timeout=ssl_handshake_timeout, + ssl_shutdown_timeout=ssl_shutdown_timeout ) yield client finally: if transport is not None: @@ -735,24 +799,34 @@ async def SocketPair( if getattr(socket, "AF_UNIX", None) is None: pytest.skip("SocketPair requires socket.AF_UNIX and is not supported on current platform") - if ct.name == "unix": - sock, peer = socket.socketpair(socket.AF_UNIX, socket.SOCK_STREAM) - elif ct.name == "udp": - sock, peer = socket.socketpair(socket.AF_UNIX, socket.SOCK_DGRAM) + if ct.name == "udp": + sock, peer = socket.socketpair(socket.AF_UNIX, socket.SOCK_DGRAM) else: - pytest.skip(f"SocketPair is not supported for {ct.name}") + sock, peer = socket.socketpair(socket.AF_UNIX, socket.SOCK_STREAM) try: - async with TestClient(ct=ct, sock=sock, is_buffered=server_is_buffered, - protocol_factory=server_protocol_factory, - ssl_handshake_timeout=server_ssl_handshake_timeout, - ssl_shutdown_timeout=server_ssl_shutdown_timeout) as server: - async with TestClient(ct=ct, sock=peer, - server_hostname=client_server_hostname, - is_buffered=client_is_buffered, - protocol_factory=client_protocol_factory, - ) as client: - yield server, client + async with AsyncExitStack() as stack: + server_context = TestClient( + ct=ct, + is_buffered=server_is_buffered, + protocol_factory=server_protocol_factory, + ssl_handshake_timeout=server_ssl_handshake_timeout, + ssl_shutdown_timeout=server_ssl_shutdown_timeout, + sock=sock, + sock_server_side=True, + ) + client_context = TestClient( + ct=ct, + sock=peer, + server_hostname=client_server_hostname, + is_buffered=client_is_buffered, + protocol_factory=client_protocol_factory, + ) + server, client = await asyncio.gather( + stack.enter_async_context(server_context), + stack.enter_async_context(client_context), + ) + yield server, client finally: sock.close() peer.close()