From cb1dc2c2e56f8c773f29663b73af17629b3cd877 Mon Sep 17 00:00:00 2001 From: Elvis Pranskevichus Date: Sat, 19 Sep 2026 19:55:35 -0700 Subject: [PATCH] Align SSL startup retries with libpq and fix rejection diagnostics An `ErrorResponse` to `SSLRequest` is currently reported as an SSL refusal and transport fallback only handles a narrow set of authorization errors. Follow libpq transport-selection behavior and improve the SSL exchange diagnostics. Be careful around trusting pre-TLS server text or exposing SQLSTATE (see CVE-2024-10977). Make it so that alternative transports in allow/prefer retries are only tried _before_ authentication succeeds, not after. While here, align asyncpg with libpq and require `ssl` to be set to a mandatory TLS mode (or explicit context) when `direct_tls` is specified. Fixes #1317. Closes #1346. Closes #1348. --- asyncpg/_testbase/__init__.py | 2 +- asyncpg/_testbase/fuzzer.py | 103 +++- asyncpg/connect_utils.py | 99 ++-- asyncpg/connection.py | 18 +- asyncpg/protocol/coreproto.pxd | 1 + asyncpg/protocol/coreproto.pyx | 2 + asyncpg/protocol/protocol.pyi | 1 + asyncpg/protocol/protocol.pyx | 1 + tests/test_connect.py | 826 ++++++++++++++++++++++++--------- 9 files changed, 782 insertions(+), 271 deletions(-) diff --git a/asyncpg/_testbase/__init__.py b/asyncpg/_testbase/__init__.py index 95775e11..16238ce1 100644 --- a/asyncpg/_testbase/__init__.py +++ b/asyncpg/_testbase/__init__.py @@ -402,7 +402,7 @@ def setUpClass(cls): host = '127.0.0.1' cls.proxy = fuzzer.TCPFuzzingProxy( backend_host=host, - backend_port=conn_spec['port'], + backend_port=int(conn_spec['port']), ) cls.proxy.start() diff --git a/asyncpg/_testbase/fuzzer.py b/asyncpg/_testbase/fuzzer.py index 88745646..99ad3667 100644 --- a/asyncpg/_testbase/fuzzer.py +++ b/asyncpg/_testbase/fuzzer.py @@ -6,6 +6,7 @@ import asyncio +import contextlib import socket import threading import typing @@ -34,6 +35,7 @@ def __init__(self, *, listening_addr: str='127.0.0.1', self.connections = {} self.sock = None self.listen_task = None + self.connection_factory = Connection async def _wait(self, work): work_task = asyncio.ensure_future(work) @@ -98,11 +100,13 @@ async def _main(self, started_event): try: await self.listen_task finally: + tasks = list(self.connections.values()) for c in list(self.connections): c.close() - await asyncio.sleep(0.01) + await asyncio.gather(*tasks, return_exceptions=True) if hasattr(self.loop, 'remove_reader'): - self.loop.remove_reader(self.sock.fileno()) + with contextlib.suppress(NotImplementedError): + self.loop.remove_reader(self.sock.fileno()) self.sock.close() async def listen(self): @@ -119,7 +123,7 @@ async def listen(self): except StopServer: break - conn = Connection(client_sock, backend_sock, self) + conn = self.connection_factory(client_sock, backend_sock, self) conn_task = self.loop.create_task(conn.handle()) self.connections[conn] = conn_task @@ -141,13 +145,13 @@ def reset(self): self.restore_connectivity() def _close_connection(self, connection): - conn_task = self.connections.pop(connection, None) + conn_task = self.connections.get(connection) if conn_task is not None: conn_task.cancel() def close_all_connections(self): for conn in list(self.connections): - self.loop.call_soon_threadsafe(self._close_connection, conn) + self.loop.call_soon_threadsafe(conn.close) class Connection: @@ -161,6 +165,27 @@ def __init__(self, client_sock, backend_sock, proxy): self.proxy_to_backend_task = None self.proxy_from_backend_task = None self.is_closed = False + self.client_reader = None + self.client_writer = None + + async def prepare(self): + """Optionally negotiate or inject faults before forwarding bytes.""" + return True + + async def use_client_stream(self, ssl_context=None): + """Wrap the accepted socket in a stream, optionally with direct TLS.""" + self.client_reader = asyncio.StreamReader() + protocol = asyncio.StreamReaderProtocol(self.client_reader) + transport, _ = await self.loop.connect_accepted_socket( + lambda: protocol, self.client_sock, ssl=ssl_context) + self.client_writer = asyncio.StreamWriter( + transport, protocol, self.client_reader, self.loop) + + async def start_tls(self, ssl_context): + writer = self.client_writer + writer._transport = await self.loop.start_tls( + writer.transport, writer._protocol, ssl_context, server_side=True) + writer._protocol._over_ssl = True def close(self): if self.is_closed: @@ -170,48 +195,72 @@ def close(self): if self.proxy_to_backend_task is not None: self.proxy_to_backend_task.cancel() - self.proxy_to_backend_task = None if self.proxy_from_backend_task is not None: self.proxy_from_backend_task.cancel() - self.proxy_from_backend_task = None self.proxy._close_connection(self) async def handle(self): - self.proxy_to_backend_task = asyncio.ensure_future( - self.proxy_to_backend()) + try: + if not await self.prepare(): + return - self.proxy_from_backend_task = asyncio.ensure_future( - self.proxy_from_backend()) + self.proxy_to_backend_task = asyncio.ensure_future( + self.proxy_to_backend()) + self.proxy_from_backend_task = asyncio.ensure_future( + self.proxy_from_backend()) - try: await asyncio.wait( [self.proxy_to_backend_task, self.proxy_from_backend_task], return_when=asyncio.FIRST_COMPLETED) + except ConnectionError: + pass + finally: + # Relay completion can schedule close() while cleanup is awaiting + # its children. Do not let that callback cancel cleanup itself. + self.is_closed = True if self.proxy_to_backend_task is not None: self.proxy_to_backend_task.cancel() if self.proxy_from_backend_task is not None: self.proxy_from_backend_task.cancel() + await asyncio.gather(*( + task for task in (self.proxy_to_backend_task, + self.proxy_from_backend_task) + if task is not None + ), return_exceptions=True) + # Asyncio fails to properly remove the readers and writers # when the task doing recv() or send() is cancelled, so # we must remove the readers and writers manually before # closing the sockets. - self.loop.remove_reader(self.client_sock.fileno()) - self.loop.remove_writer(self.client_sock.fileno()) - self.loop.remove_reader(self.backend_sock.fileno()) - self.loop.remove_writer(self.backend_sock.fileno()) - - self.client_sock.close() - self.backend_sock.close() + sockets = [self.backend_sock] + if self.client_writer is None: + sockets.append(self.client_sock) + else: + self.client_writer.close() + with contextlib.suppress(ConnectionError): + await self.client_writer.wait_closed() + for sock in sockets: + if sock.fileno() < 0: + continue + # ProactorEventLoop has no reader/writer registration API. + with contextlib.suppress(NotImplementedError): + self.loop.remove_reader(sock.fileno()) + self.loop.remove_writer(sock.fileno()) + sock.close() + self.proxy.connections.pop(self, None) async def _read(self, sock, n): - read_task = asyncio.ensure_future( - self.loop.sock_recv(sock, n)) + if sock is self.client_sock and self.client_reader is not None: + read = self.client_reader.read(n) + else: + read = self.loop.sock_recv(sock, n) + read_task = asyncio.ensure_future(read) conn_event_task = asyncio.ensure_future( self.connectivity_loss.wait()) @@ -230,10 +279,16 @@ async def _read(self, sock, n): read_task.cancel() if not conn_event_task.done(): conn_event_task.cancel() + await asyncio.gather(read_task, conn_event_task, + return_exceptions=True) async def _write(self, sock, data): - write_task = asyncio.ensure_future( - self.loop.sock_sendall(sock, data)) + if sock is self.client_sock and self.client_writer is not None: + self.client_writer.write(data) + write = self.client_writer.drain() + else: + write = self.loop.sock_sendall(sock, data) + write_task = asyncio.ensure_future(write) conn_event_task = asyncio.ensure_future( self.connectivity_loss.wait()) @@ -252,6 +307,8 @@ async def _write(self, sock, data): write_task.cancel() if not conn_event_task.done(): conn_event_task.cancel() + await asyncio.gather(write_task, conn_event_task, + return_exceptions=True) async def proxy_to_backend(self): buf = None diff --git a/asyncpg/connect_utils.py b/asyncpg/connect_utils.py index 651bff58..c3a8fe8a 100644 --- a/asyncpg/connect_utils.py +++ b/asyncpg/connect_utils.py @@ -33,6 +33,9 @@ from . import protocol +_SSL_REQUEST_CODE = 80877103 + + class SSLMode(enum.IntEnum): disable = 0 allow = 1 @@ -811,9 +814,15 @@ def _parse_connect_dsn_and_args(*, dsn, host, port, user, elif ssl is True: ssl = ssl_module.create_default_context() sslmode = SSLMode.verify_full + elif isinstance(ssl, ssl_module.SSLContext): + sslmode = SSLMode.require else: sslmode = SSLMode.disable + if sslneg is SSLNegotiation.direct and sslmode < SSLMode.require: + raise exceptions.ClientConfigurationError( + 'direct TLS requires sslmode=require, verify-ca, or verify-full') + if server_settings is not None and ( not isinstance(server_settings, dict) or not all(isinstance(k, str) for k in server_settings) or @@ -924,6 +933,8 @@ def __init__( self.ssl_is_advisory = ssl_is_advisory def data_received(self, data: bytes) -> None: + if self.on_data.done(): + return if data == b'S': self.on_data.set_result(True) elif (self.ssl_is_advisory and @@ -934,6 +945,13 @@ def data_received(self, data: bytes) -> None: # sslmode=prefer. But be extra sure to disallow insecure # connections when the ssl context asks for real security. self.on_data.set_result(False) + elif data.startswith(b'E'): + # Never trust or expose pre-TLS ErrorResponse fields + # (including SQLSTATE). See CVE-2024-10977. + self.on_data.set_exception(ConnectionError( + f'PostgreSQL server at "{self.host}:{self.port}": ' + 'server sent an error response during SSL exchange; ' + 'check the server logs for details')) else: self.on_data.set_exception( ConnectionError( @@ -964,6 +982,7 @@ async def _create_ssl_connection( loop: asyncio.AbstractEventLoop, ssl_context: ssl_module.SSLContext, ssl_is_advisory: bool = False, + retry: bool = False, ) -> typing.Tuple[asyncio.Transport, _ProctolFactoryR]: tr, pr = await loop.create_connection( @@ -971,7 +990,7 @@ async def _create_ssl_connection( ssl_context, ssl_is_advisory), host, port) - tr.write(struct.pack('!ll', 8, 80877103)) # SSLRequest message. + tr.write(struct.pack('!ll', 8, _SSL_REQUEST_CODE)) # SSLRequest message. try: do_ssl_upgrade = await pr.on_data @@ -984,16 +1003,22 @@ async def _create_ssl_connection( try: new_tr = await loop.start_tls( tr, pr, ssl_context, server_hostname=host) - assert new_tr is not None - except (Exception, asyncio.CancelledError): + if new_tr is None: + raise ConnectionError('connection closed during TLS ' + 'handshake') + except (Exception, asyncio.CancelledError) as exc: tr.close() + if retry and isinstance(exc, OSError) and not isinstance( + exc, asyncio.TimeoutError + ): + raise _RetryConnectSignal() from exc raise else: new_tr = tr pg_proto = protocol_factory() - pg_proto.is_ssl = do_ssl_upgrade pg_proto.connection_made(new_tr) + pg_proto.is_ssl = do_ssl_upgrade new_tr.set_protocol(pg_proto) return new_tr, pg_proto @@ -1014,8 +1039,11 @@ async def _create_ssl_connection( new_tr, pg_proto = await conn_factory(sock=sock) pg_proto.is_ssl = do_ssl_upgrade return new_tr, pg_proto - except (Exception, asyncio.CancelledError): + except (Exception, asyncio.CancelledError) as exc: sock.close() + if (retry and do_ssl_upgrade and isinstance(exc, OSError) + and not isinstance(exc, asyncio.TimeoutError)): + raise _RetryConnectSignal() from exc raise @@ -1040,7 +1068,9 @@ async def _connect_addr( args = (addr, loop, config, connection_class, record_class, params_input) # prepare the params (which attempt has ssl) for the 2 attempts - if params.sslmode == SSLMode.allow: + if isinstance(addr, str): + return await __connect_addr(params, False, *args) + elif params.sslmode == SSLMode.allow: params_retry = params params = params._replace(ssl=None) elif params.sslmode == SSLMode.prefer: @@ -1059,24 +1089,33 @@ async def _connect_addr( # second attempt try: return await __connect_addr(params_retry, False, *args) - except ( - exceptions.InvalidAuthorizationSpecificationError, - exceptions.ConnectionDoesNotExistError, - ) as second_error: + except _NextHostSignal as exc: + # The server error will be unwrapped by _connect on host exhaustion. + exc.__cause__.__cause__ = first_error + raise + except (Exception, asyncio.CancelledError) as second_error: # If the preferred attempt produced a useful authentication error, # do not hide it behind a generic rejection from the fallback mode. if ( isinstance(first_error, exceptions.InvalidPasswordError) and not isinstance(second_error, exceptions.InvalidPasswordError) + and isinstance(second_error, ( + exceptions.InvalidAuthorizationSpecificationError, + exceptions.ConnectionDoesNotExistError, + )) ): raise first_error from None - raise + raise second_error from first_error class _RetryConnectSignal(Exception): pass +class _NextHostSignal(Exception): + pass + + async def __connect_addr( params, retry, @@ -1106,7 +1145,8 @@ async def __connect_addr( elif params.ssl: connector = _create_ssl_connection( proto_factory, *addr, loop=loop, ssl_context=params.ssl, - ssl_is_advisory=params.sslmode == SSLMode.prefer) + ssl_is_advisory=params.sslmode == SSLMode.prefer, + retry=retry and params.sslmode == SSLMode.prefer) else: connector = loop.create_connection(proto_factory, *addr) @@ -1114,34 +1154,23 @@ async def __connect_addr( try: await connected - except ( - exceptions.InvalidAuthorizationSpecificationError, - exceptions.ConnectionDoesNotExistError, # seen on Windows - ) as exc: + except exceptions.PostgresError as exc: tr.close() - # retry=True here is a redundant check because we don't want to - # accidentally raise the internal _RetryConnectSignal to the user - if retry and ( + if not pr._auth_received and isinstance( + exc, exceptions.CannotConnectNowError + ): + raise _NextHostSignal() from exc + + # Only retry before AuthenticationOk, and only if the other transport + # has not been tried. This includes ConnectionDoesNotExistError, + # which is needed for Windows compatibility. + if retry and not pr._auth_received and ( params.sslmode == SSLMode.allow and not pr.is_ssl or params.sslmode == SSLMode.prefer and pr.is_ssl ): - # Trigger retry when: - # 1. First attempt with sslmode=allow, ssl=None failed - # 2. First attempt with sslmode=prefer, ssl=ctx failed while the - # server claimed to support SSL (returning "S" for SSLRequest) - # (likely because pg_hba.conf rejected the connection) raise _RetryConnectSignal() from exc - - else: - # but will NOT retry if: - # 1. First attempt with sslmode=prefer failed but the server - # doesn't support SSL (returning 'N' for SSLRequest), because - # we already tried to connect without SSL thru ssl_is_advisory - # 2. Second attempt with sslmode=prefer, ssl=None failed - # 3. Second attempt with sslmode=allow, ssl=ctx failed - # 4. Any other sslmode - raise + raise except (Exception, asyncio.CancelledError): tr.close() @@ -1243,6 +1272,8 @@ async def _connect(*, loop, connection_class, record_class, **kwargs): break except OSError as ex: last_error = ex + except _NextHostSignal as ex: + last_error = ex.__cause__ else: if target_attr == SessionAttribute.prefer_standby and candidates: chosen_connection = random.choice(candidates) diff --git a/asyncpg/connection.py b/asyncpg/connection.py index 71fb04f8..a0010d0b 100644 --- a/asyncpg/connection.py +++ b/asyncpg/connection.py @@ -2251,6 +2251,15 @@ async def connect(dsn=None, *, The default is ``'prefer'``: try an SSL connection and fallback to non-SSL connection if that fails. + With ``'allow'`` and ``'prefer'``, a server error before + ``AuthenticationOk`` permits one retry using the other transport. + + Errors in response to SSLRequest, timeouts, cancellation, client-side + authentication errors, and errors after ``AuthenticationOk`` do not + trigger transport retries. A server reporting that it cannot accept + connections yet (SQLSTATE ``57P03``) before ``AuthenticationOk`` causes + asyncpg to try the next host. + .. note:: *ssl* is ignored for Unix domain socket communication. @@ -2302,7 +2311,8 @@ async def connect(dsn=None, *, :param bool direct_tls: Pass ``True`` to skip PostgreSQL STARTTLS mode and perform a direct - SSL connection. Must be used alongside ``ssl`` param. + SSL connection. Requires ``ssl='require'``, ``'verify-ca'``, + ``'verify-full'``, ``True``, or an explicit ``SSLContext``. :param dict server_settings: An optional dict of server runtime parameters. Refer to @@ -2417,6 +2427,12 @@ async def connect(dsn=None, *, .. versionchanged:: 0.31.0 Added the *servicefile* and *service* parameters. + .. versionchanged:: 0.32.0 + ``direct_tls=True`` requires an SSL mode of ``'require'`` or higher, + ``ssl=True``, or an explicit ``SSLContext``. Other values + (``'disable'``, ``'allow'``, and ``'prefer'``) will raise a + ``ClientConfigurationError``. + .. _SSLContext: https://docs.python.org/3/library/ssl.html#ssl.SSLContext .. _create_default_context: https://docs.python.org/3/library/ssl.html#ssl.create_default_context diff --git a/asyncpg/protocol/coreproto.pxd b/asyncpg/protocol/coreproto.pxd index 34c7c712..1cfd9e90 100644 --- a/asyncpg/protocol/coreproto.pxd +++ b/asyncpg/protocol/coreproto.pxd @@ -79,6 +79,7 @@ cdef class CoreProtocol: str _execute_stmt_name ConnectionStatus con_status + readonly bint _auth_received ProtocolState state TransactionStatus xact_status diff --git a/asyncpg/protocol/coreproto.pyx b/asyncpg/protocol/coreproto.pyx index eae753df..c9c148dd 100644 --- a/asyncpg/protocol/coreproto.pyx +++ b/asyncpg/protocol/coreproto.pyx @@ -32,6 +32,7 @@ cdef class CoreProtocol: self.auth_msg = None self.con_params = con_params self.con_status = CONNECTION_BAD + self._auth_received = False self.state = PROTOCOL_IDLE self.xact_status = PQTRANS_IDLE self.encoding = 'utf-8' @@ -573,6 +574,7 @@ cdef class CoreProtocol: if status == AUTH_SUCCESSFUL: # AuthenticationOk + self._auth_received = True self.result_type = RESULT_OK elif status == AUTH_REQUIRED_PASSWORD: diff --git a/asyncpg/protocol/protocol.pyi b/asyncpg/protocol/protocol.pyi index 34db6440..3bed28df 100644 --- a/asyncpg/protocol/protocol.pyi +++ b/asyncpg/protocol/protocol.pyi @@ -104,6 +104,7 @@ class PreparedStatementState(Generic[_Record]): def __reduce__(self) -> Any: ... class CoreProtocol: + _auth_received: bool backend_pid: Any backend_secret: Any __pyx_vtable__: Any diff --git a/asyncpg/protocol/protocol.pyx b/asyncpg/protocol/protocol.pyx index 91735c87..e7627d55 100644 --- a/asyncpg/protocol/protocol.pyx +++ b/asyncpg/protocol/protocol.pyx @@ -969,6 +969,7 @@ cdef class BaseProtocol(CoreProtocol): def connection_made(self, transport): self.transport = transport + self._is_ssl = transport.get_extra_info('ssl_object') is not None sock = transport.get_extra_info('socket') if (sock is not None and diff --git a/tests/test_connect.py b/tests/test_connect.py index fae32447..839c3dcf 100644 --- a/tests/test_connect.py +++ b/tests/test_connect.py @@ -7,8 +7,10 @@ import asyncio import contextlib +import functools import gc import ipaddress +import itertools import os import pathlib import platform @@ -16,6 +18,7 @@ import socket import ssl import stat +import struct import tempfile import textwrap import unittest @@ -28,6 +31,7 @@ import asyncpg from asyncpg import _testbase as tb +from asyncpg._testbase import fuzzer from asyncpg import connection as pg_connection from asyncpg import connect_utils from asyncpg import cluster as pg_cluster @@ -606,7 +610,8 @@ class TestConnectParams(tb.TestCase): 'PGSSLNEGOTIATION': 'postgres' }, - 'dsn': 'postgres://u:p@localhost/d?sslnegotiation=direct', + 'dsn': 'postgres://u:p@localhost/d' + '?sslnegotiation=direct&sslmode=require', 'result': ([('localhost', 5432)], { 'user': 'u', @@ -620,7 +625,8 @@ class TestConnectParams(tb.TestCase): { 'name': 'params_ssl_negotiation_env', 'env': { - 'PGSSLNEGOTIATION': 'direct' + 'PGSSLNEGOTIATION': 'direct', + 'PGSSLMODE': 'require', }, 'dsn': 'postgres://u:p@localhost/d', @@ -1648,12 +1654,14 @@ async def test_connect_args_validation(self): with self.assertRaisesRegex(ValueError, 'greater than 0'): await asyncpg.connect(command_timeout=val) - for arg in {'max_cacheable_statement_size', - 'max_cached_statement_lifetime', - 'statement_cache_size'}: - for val in {None, -1, True, False}: - with self.assertRaisesRegex(ValueError, 'greater or equal'): - await asyncpg.connect(**{arg: val}) + cases = itertools.product( + ('max_cacheable_statement_size', + 'max_cached_statement_lifetime', 'statement_cache_size'), + (None, -1, True, False), + ) + for arg, val in cases: + with self.assertRaisesRegex(ValueError, 'greater or equal'): + await asyncpg.connect(**{arg: val}) class TestConnection(tb.ConnectedTestCase): @@ -1793,7 +1801,117 @@ async def test_connection_no_home_dir(self): ssl='verify-full') -class BaseTestSSLConnection(tb.ConnectedTestCase): +class SSLProxyConnection(fuzzer.Connection): + """PostgreSQL startup faults using the shared proxy's socket lifecycle.""" + + # Startup packets have no message-type byte: length, then protocol code. + STARTUP_HEADER = struct.Struct('!II') + SSL_REQUEST = STARTUP_HEADER.pack( + STARTUP_HEADER.size, connect_utils._SSL_REQUEST_CODE) + + # Subsequent messages have a type byte and a length including this field. + MESSAGE_LENGTH = struct.Struct('!I') + AUTHENTICATION_OK = 0 + TLS_RECORD_HEADER_SIZE = 5 + + def __init__(self, *args, faults, direct, attempts, tasks, caller_loop): + super().__init__(*args) + self.faults = faults + self.direct = direct + self.attempts = attempts + self.tasks = tasks + self.caller_loop = caller_loop + + async def prepare(self): + """Return True to forward traffic, False to close a faulted attempt.""" + self.tasks.append(asyncio.current_task()) + index = len(self.attempts) + self.attempts.append(None) + fault = self.faults[index] if index < len(self.faults) else {} + + ssl_context = ssl.SSLContext(ssl.PROTOCOL_TLS_SERVER) + ssl_context.load_cert_chain(SSL_CERT_FILE, SSL_KEY_FILE) + await self.use_client_stream(ssl_context if self.direct else None) + header = await self.client_reader.readexactly(self.STARTUP_HEADER.size) + ssl_requested = header == self.SSL_REQUEST + self.attempts[index] = ssl_requested or self.direct + if not fault: + await self.loop.sock_sendall(self.backend_sock, header) + return True + + if ssl_requested: + if not await self._negotiate_ssl(fault, ssl_context): + return False + header = await self.client_reader.readexactly( + self.STARTUP_HEADER.size) + + # Consume the startup parameters before injecting an authentication + # failure, disconnect, or stall. + length, _ = self.STARTUP_HEADER.unpack(header) + await self.client_reader.readexactly(length - self.STARTUP_HEADER.size) + if fault.get('hang') == 'startup': + await self._wait_for_client_close(ready=fault['ready']) + elif not fault.get('disconnect'): + if fault.get('auth_ok'): + self._write_message( + b'R', self.MESSAGE_LENGTH.pack(self.AUTHENTICATION_OK)) + self._write_error_response(fault['error']) + await self.client_writer.drain() + await self._wait_for_client_close() + return False + + async def _negotiate_ssl(self, fault, ssl_context): + """Return whether the client can proceed to its startup message.""" + if fault.get('hang') == 'sslrequest': + await self._wait_for_client_close(ready=fault['ready']) + return False + + response = fault.get('response', b'S') + self.client_writer.write(response) + await self.client_writer.drain() + if response == b'S': + if fault.get('hang') == 'handshake': + # Wait until ClientHello starts before signalling readiness. + await self.client_reader.readexactly( + self.TLS_RECORD_HEADER_SIZE) + await self._wait_for_client_close(ready=fault['ready']) + return False + await self.start_tls(ssl_context) + elif response != b'N': + # Do not send the rest of a fragmented ErrorResponse: the + # diagnostic must neither wait for nor expose its payload. + await self._wait_for_client_close() + return False + return True + + async def _wait_for_client_close(self, *, ready=None): + if ready is not None: + # Readiness events belong to the test loop, not the proxy thread. + self.caller_loop.call_soon_threadsafe(ready.set) + await self.client_reader.read() + + def _write_message(self, message_type, payload): + length = self.MESSAGE_LENGTH.size + len(payload) + self.client_writer.write( + message_type + self.MESSAGE_LENGTH.pack(length) + payload) + + def _write_error_response(self, exc_type): + fields = ( + (b'S', b'FATAL'), # Severity + (b'C', exc_type.sqlstate.encode('ascii')), # SQLSTATE + (b'M', b'proxy startup error'), # Message + ) + payload = b''.join(tag + value + b'\0' for tag, value in fields) + self._write_message(b'E', payload + b'\0') + + +class BaseTestSSLConnection(tb.ConnectedTestCase, tb.ProxiedClusterTestCase): + @classmethod + def get_connection_spec(cls, kwargs={}): + # Keep administrative connections and unproxied SSL tests on the + # cluster. Proxy tests explicitly supply the shared proxy's port. + return tb.ClusterTestCase.get_connection_spec.__func__(cls, kwargs) + @classmethod def get_server_settings(cls): conf = super().get_server_settings() @@ -1847,19 +1965,330 @@ def tearDown(self): def _add_hba_entry(self): raise NotImplementedError() + def _add_hba_entries(self, connection_type, auth_method, *, + addresses=('127.0.0.0/24', '::1/128')): + for address in addresses: + self.cluster.add_hba_entry( + type=connection_type, address=ipaddress.ip_network(address), + database='postgres', user='ssl_user', auth_method=auth_method) + + async def _test_works(self, *, expected_ssl=None, **conn_args): + con = await self.connect(**conn_args) + try: + self.assertEqual(await con.fetchval('SELECT 42'), 42) + if expected_ssl is not None: + self.assertEqual(con._protocol.is_ssl, expected_ssl) + finally: + await con.close() + + async def _test_cancellation_recovery(self, con): + self.assertEqual(await con.fetchval('SELECT 42'), 42) + with self.assertRaises(asyncio.TimeoutError): + await con.execute('SELECT pg_sleep(5)', timeout=0.5) + self.assertEqual(await con.fetchval('SELECT 43'), 43) + + async def _test_pool(self, *, expected_ssl, **conn_args): + pool = await self.create_pool(min_size=5, max_size=10, **conn_args) + + async def worker(): + async with pool.acquire() as con: + self.assertEqual(con._protocol.is_ssl, expected_ssl) + await self._test_cancellation_recovery(con) + + tasks = [self.loop.create_task(worker()) for _ in range(100)] + try: + await asyncio.gather(*tasks) + finally: + for task in tasks: + task.cancel() + await asyncio.gather(*tasks, return_exceptions=True) + await pool.close() + + async def _test_sslmode_works( + self, sslmode, *, expected_ssl, host='localhost', + ): + await self._test_works( + dsn='postgresql://foo/postgres?sslmode=' + sslmode, + host=host, user='ssl_user', expected_ssl=expected_ssl) + + async def _test_sslmode_fails( + self, sslmode, *, host='localhost', + exc_type=asyncpg.InvalidAuthorizationSpecificationError, + ): + # XXX: uvloop artifact + old_handler = self.loop.get_exception_handler() + try: + self.loop.set_exception_handler(lambda *args: None) + with self.assertRaises(exc_type): + await self._test_works( + dsn='postgresql://foo/?sslmode=' + sslmode, + host=host, user='ssl_user') + finally: + self.loop.set_exception_handler(old_handler) + + async def _test_connect_interruption(self, *, ready, cancel, **conn_args): + async def connect(): + # Python 3.9 wraps CancelledError when it crosses a task boundary. + # Capture the original exception so its cause can be checked. + expected = (asyncio.CancelledError if cancel + else asyncio.TimeoutError) + with self.assertRaises(expected) as raised: + await self.connect(timeout=2 if cancel else 0.2, **conn_args) + return raised.exception + + task = self.loop.create_task(connect()) + try: + if cancel: + await asyncio.wait_for(ready.wait(), 2) + task.cancel() + return await task + finally: + if not task.done(): + task.cancel() + await asyncio.gather(task, return_exceptions=True) + + @contextlib.asynccontextmanager + async def connection_proxy(self, *faults, direct=False): + attempts = [] + tasks = [] + caller_loop = asyncio.get_running_loop() + + async def configure(): + self.proxy.connection_factory = functools.partial( + SSLProxyConnection, faults=faults, direct=direct, + attempts=attempts, tasks=tasks, caller_loop=caller_loop) + + async def finish(): + try: + # Discarded transports must close without forced cleanup. + results = await asyncio.wait_for(asyncio.gather( + *tasks, return_exceptions=True), 2) + for result in results: + if isinstance(result, BaseException) and not isinstance( + result, asyncio.CancelledError + ): + raise result + finally: + self.proxy.connection_factory = fuzzer.Connection + + await asyncio.wrap_future(asyncio.run_coroutine_threadsafe( + configure(), self.proxy.loop)) + try: + yield self.proxy.listening_port, attempts + finally: + await asyncio.wrap_future(asyncio.run_coroutine_threadsafe( + finish(), self.proxy.loop)) + @unittest.skipIf(os.environ.get('PGHOST'), 'unmanaged cluster') class TestSSLConnection(BaseTestSSLConnection): def _add_hba_entry(self): - self.cluster.add_hba_entry( - type='hostssl', address=ipaddress.ip_network('127.0.0.0/24'), - database='postgres', user='ssl_user', - auth_method='trust') + self._add_hba_entries('hostssl', 'trust') - self.cluster.add_hba_entry( - type='hostssl', address=ipaddress.ip_network('::1/128'), - database='postgres', user='ssl_user', - auth_method='trust') + def _allow_plaintext(self): + self._add_hba_entries('hostnossl', 'trust', + addresses=('127.0.0.0/24',)) + self.cluster.reload() + + async def test_ssl_startup_retry_boundary(self): + self._allow_plaintext() + cases = itertools.product( + ('allow', 'prefer'), + (False, True), + (asyncpg.InvalidAuthorizationSpecificationError, + asyncpg.InvalidCatalogNameError, + asyncpg.TooManyConnectionsError, + asyncpg.ConnectionDoesNotExistError), + ) + for mode, auth_ok, error in cases: + with self.subTest(mode=mode, auth_ok=auth_ok, error=error): + fault = {'auth_ok': auth_ok, 'error': error} + async with self.connection_proxy(fault) as result: + port, seen = result + calls = [] + + async def password(): + calls.append(True) + return 'unused with trust authentication' + + options = dict(host='127.0.0.1', port=port, + user='ssl_user', ssl=mode, + password=password) + if auth_ok: + with self.assertRaises(error): + await self.connect(**options) + self.assertEqual(seen, [mode == 'prefer']) + else: + await self._test_works( + expected_ssl=mode == 'allow', **options) + self.assertEqual( + seen, [mode == 'prefer', mode == 'allow']) + self.assertEqual(len(calls), 1) + + async def test_ssl_post_authentication_server_error(self): + # PostgreSQL processes startup GUCs after AuthenticationOk. + self._allow_plaintext() + for mode in ('allow', 'prefer'): + async with self.connection_proxy() as (port, seen): + with self.assertRaises(asyncpg.InvalidParameterValueError): + await self.connect( + host='127.0.0.1', port=port, user='ssl_user', ssl=mode, + server_settings={'work_mem': 'invalid'}) + self.assertEqual(seen, [mode == 'prefer']) + + async def test_ssl_fallback_error_cause(self): + for mode in ('allow', 'prefer'): + faults = ( + {'error': asyncpg.InvalidCatalogNameError}, + {'error': asyncpg.TooManyConnectionsError}, + ) + async with self.connection_proxy(*faults) as (port, seen): + with self.assertRaises(asyncpg.TooManyConnectionsError) as exc: + await self.connect(host='127.0.0.1', port=port, + user='ssl_user', ssl=mode) + self.assertIsInstance(exc.exception.__cause__, + asyncpg.InvalidCatalogNameError) + self.assertEqual(seen, [mode == 'prefer', mode == 'allow']) + + async def test_ssl_cannot_connect_now_failover(self): + self._allow_plaintext() + cases = itertools.product( + ('allow', 'prefer', 'require', 'disable'), (False, True)) + for mode, exhausted in cases: + with self.subTest(mode=mode, exhausted=exhausted): + faults = [{'error': asyncpg.CannotConnectNowError}] + if exhausted: + faults.append({'error': asyncpg.CannotConnectNowError}) + async with self.connection_proxy(*faults) as (port, seen): + options = dict(host=['127.0.0.1'] * 2, + port=[port, port], + user='ssl_user', ssl=mode) + if exhausted: + with self.assertRaises( + asyncpg.CannotConnectNowError + ): + await self.connect(**options) + else: + await self._test_works(**options) + # Both hosts start with the same transport; a fallback + # on the first host would use the opposite transport. + self.assertEqual(seen, [mode in ('prefer', 'require')] + * 2) + for mode in ('allow', 'prefer'): + async with self.connection_proxy( + {'error': asyncpg.InvalidCatalogNameError}, + {'error': asyncpg.CannotConnectNowError}, + ) as (port, seen): + with self.assertRaises( + asyncpg.CannotConnectNowError + ) as raised: + await self.connect(host='127.0.0.1', port=port, + user='ssl_user', ssl=mode) + self.assertIsInstance(raised.exception.__cause__, + asyncpg.InvalidCatalogNameError) + self.assertEqual(len(seen), 2) + async with self.connection_proxy( + {'auth_ok': True, 'error': asyncpg.CannotConnectNowError} + ) as (port, seen): + with self.assertRaises(asyncpg.CannotConnectNowError): + await self.connect(host=['127.0.0.1'] * 2, + port=[port, port], user='ssl_user', + ssl=mode) + self.assertEqual(len(seen), 1) + + async def test_sslrequest_errors_are_terminal(self): + # E alone tests fragmented ErrorResponse delivery: diagnostics must + # not wait for or expose the rest of the unauthenticated message. + error_response = (b'E\0\0\0\xffC' + + asyncpg.CannotConnectNowError.sqlstate.encode() + + b'\0Muntrusted\0\0') + for response in (b'E', error_response, b'X', b'Sextra', b'Nextra'): + with self.subTest(response=response): + async with self.connection_proxy( + {'response': response} + ) as (port, seen): + with self.assertRaises(ConnectionError) as raised: + await self.connect(host='127.0.0.1', port=port, + user='ssl_user', ssl='prefer') + if response.startswith(b'E'): + self.assertEqual( + str(raised.exception), + f'PostgreSQL server at "127.0.0.1:{port}": ' + 'server sent an error response during SSL ' + 'exchange; check the server logs for details') + self.assertEqual(seen, [True]) + + async def test_ssl_declined_and_disconnect_attempt_counts(self): + for mode in ('allow', 'prefer'): + self._allow_plaintext() + async with self.connection_proxy( + {'disconnect': True} + ) as (port, seen): + await self._test_works( + host='127.0.0.1', port=port, user='ssl_user', ssl=mode) + self.assertEqual(seen, [mode == 'prefer', mode == 'allow']) + async with self.connection_proxy( + {'response': b'N', + 'error': asyncpg.InvalidAuthorizationSpecificationError} + ) as (port, seen): + with self.assertRaises( + asyncpg.InvalidAuthorizationSpecificationError + ): + await self.connect(host='127.0.0.1', port=port, + user='ssl_user', ssl='prefer') + self.assertEqual(seen, [True]) + + async def test_ssl_connect_timeout_and_cancellation(self): + cases = itertools.product( + ('sslrequest', 'handshake', 'startup'), (False, True)) + for phase, cancel in cases: + with self.subTest(phase=phase, cancel=cancel): + ready = asyncio.Event() + async with self.connection_proxy( + {'hang': phase, 'ready': ready} + ) as (port, seen): + await self._test_connect_interruption( + ready=ready, cancel=cancel, host='127.0.0.1', + port=port, user='ssl_user', ssl='prefer') + self.assertEqual(seen, [True]) + + async def test_ssl_handshake_failure_cause(self): + if self.cluster.get_pg_version() < (12, 0): + self.skipTest('PostgreSQL < 12 cannot set SSL protocol version') + for mode in ('prefer', 'require'): + async with self.connection_proxy() as (port, seen): + with self.assertRaises( + asyncpg.InvalidAuthorizationSpecificationError + if mode == 'prefer' else ssl.SSLError + ) as raised: + await self.connect( + dsn='postgresql://ssl_user@127.0.0.1/postgres' + f'?sslmode={mode}' + '&ssl_min_protocol_version=TLSv1.3', + port=port) + if mode == 'prefer': + self.assertIsInstance(raised.exception.__cause__, + ssl.SSLError) + self.assertEqual(seen, [True, False] if mode == 'prefer' + else [True]) + + async def test_ssl_fallback_timeout_and_cancellation(self): + cases = itertools.product(('allow', 'prefer'), (False, True)) + for mode, cancel in cases: + with self.subTest(mode=mode, cancel=cancel): + ready = asyncio.Event() + async with self.connection_proxy( + {'error': asyncpg.InvalidCatalogNameError}, + {'hang': 'startup', 'ready': ready}, + ) as (port, seen): + error = await self._test_connect_interruption( + ready=ready, cancel=cancel, host='127.0.0.1', + port=port, user='ssl_user', ssl=mode) + if cancel: + self.assertIsInstance(error.__cause__, + asyncpg.InvalidCatalogNameError) + self.assertEqual(seen, [mode == 'prefer', + mode == 'allow']) async def test_ssl_connection_custom_context(self): ssl_context = ssl.SSLContext(ssl.PROTOCOL_TLS_CLIENT) @@ -1871,74 +2300,104 @@ async def test_ssl_connection_custom_context(self): ssl=ssl_context) try: - self.assertEqual(await con.fetchval('SELECT 42'), 42) - - with self.assertRaises(asyncio.TimeoutError): - await con.execute('SELECT pg_sleep(5)', timeout=0.5) - - self.assertEqual(await con.fetchval('SELECT 43'), 43) + await self._test_cancellation_recovery(con) finally: await con.close() - async def test_ssl_connection_sslmode(self): - async def verify_works(sslmode, *, host='localhost'): - con = None - try: - con = await self.connect( - dsn='postgresql://foo/postgres?sslmode=' + sslmode, - host=host, - user='ssl_user') - self.assertEqual(await con.fetchval('SELECT 42'), 42) - self.assertTrue(con._protocol.is_ssl) - finally: - if con: - await con.close() - - async def verify_fails(sslmode, *, host='localhost', exn_type): - # XXX: uvloop artifact - old_handler = self.loop.get_exception_handler() - con = None - try: - self.loop.set_exception_handler(lambda *args: None) - with self.assertRaises(exn_type): - con = await self.connect( - dsn='postgresql://foo/?sslmode=' + sslmode, - host=host, - user='ssl_user') - await con.fetchval('SELECT 42') - finally: - if con: - await con.close() - self.loop.set_exception_handler(old_handler) + async def test_direct_tls_connection(self): + # A TLS-terminating gateway also supports PostgreSQL versions predating + # direct TLS. Every accepted connection must execute a real query. + self._allow_plaintext() + ctx = ssl.SSLContext(ssl.PROTOCOL_TLS_CLIENT) + ctx.check_hostname = False + ctx.verify_mode = ssl.CERT_NONE + for mode in ('require', 'verify-ca', 'verify-full', True, ctx): + with self.subTest(mode=mode): + async with self.connection_proxy(direct=True) as (port, seen): + with unittest.mock.patch.dict(os.environ, { + 'SSL_CERT_FILE': SSL_CA_CERT_FILE, + }): + await self._test_works( + dsn='postgresql://ssl_user@localhost/postgres' + '?sslrootcert=' + SSL_CA_CERT_FILE, + port=port, ssl=mode, direct_tls=True, + expected_ssl=True) + self.assertEqual(seen, [True]) + self.assertEqual(ctx.verify_mode, ssl.CERT_NONE) + self.assertFalse(ctx.check_hostname) + + async def test_direct_tls_configuration_sources(self): + base_dsn = 'postgresql://ssl_user@localhost/postgres' + + async def rejected(**kwargs): + async with self.connection_proxy() as (port, seen): + with self.assertRaisesRegex( + asyncpg.ClientConfigurationError, + 'direct TLS requires sslmode', + ): + await self.connect(port=port, **kwargs) + self.assertEqual(seen, []) + + for mode in ('disable', 'allow', 'prefer'): + with self.subTest(mode=mode): + await rejected(dsn=base_dsn, ssl=mode, direct_tls=True) + await rejected(dsn=base_dsn + + f'?sslmode={mode}&sslnegotiation=direct') + with unittest.mock.patch.dict(os.environ, { + 'PGSSLMODE': mode, 'PGSSLNEGOTIATION': 'direct', + }): + await rejected(dsn=base_dsn) + # Keyword overrides both environment values. + await self._test_works( + expected_ssl=True, dsn=base_dsn, + ssl='require', direct_tls=False) + # DSN overrides both environment values. + await self._test_works( + expected_ssl=True, dsn=base_dsn + + '?sslmode=require&sslnegotiation=postgres') + with tempfile.TemporaryDirectory() as directory: + service = pathlib.Path(directory) / 'pg_service.conf' + service.write_text('[test]\nsslnegotiation=direct\n' + f'sslmode={mode}\n') + options = dict(dsn=base_dsn + '?service=test', + servicefile=str(service)) + await rejected(**options) + await self._test_works( + expected_ssl=True, **options, + ssl='require', direct_tls=False) + options['dsn'] += ('&sslmode=require' + '&sslnegotiation=postgres') + await self._test_works(expected_ssl=True, **options) + for mode in (None, False): + await rejected(dsn=base_dsn, ssl=mode, direct_tls=True) - invalid_auth_err = asyncpg.InvalidAuthorizationSpecificationError - await verify_fails('disable', exn_type=invalid_auth_err) - await verify_works('allow') - await verify_works('prefer') - await verify_works('require') - await verify_fails('verify-ca', exn_type=ValueError) - await verify_fails('verify-full', exn_type=ValueError) + async def test_ssl_connection_sslmode(self): + await self._test_sslmode_fails('disable') + await self._test_sslmode_works('allow', expected_ssl=True) + await self._test_sslmode_works('prefer', expected_ssl=True) + await self._test_sslmode_works('require', expected_ssl=True) + await self._test_sslmode_fails('verify-ca', exc_type=ValueError) + await self._test_sslmode_fails('verify-full', exc_type=ValueError) with mock_dot_postgresql(): - await verify_works('require') - await verify_works('verify-ca') - await verify_works('verify-ca', host='127.0.0.1') - await verify_works('verify-full') - await verify_fails('verify-full', host='127.0.0.1', - exn_type=ssl.CertificateError) + await self._test_sslmode_works('require', expected_ssl=True) + await self._test_sslmode_works('verify-ca', expected_ssl=True) + await self._test_sslmode_works( + 'verify-ca', host='127.0.0.1', expected_ssl=True) + await self._test_sslmode_works('verify-full', expected_ssl=True) + await self._test_sslmode_fails( + 'verify-full', host='127.0.0.1', exc_type=ssl.CertificateError) with mock_dot_postgresql(crl=True): - await verify_fails('disable', exn_type=invalid_auth_err) - await verify_works('allow') - await verify_works('prefer') - await verify_fails('require', - exn_type=ssl.SSLError) - await verify_fails('verify-ca', - exn_type=ssl.SSLError) - await verify_fails('verify-ca', host='127.0.0.1', - exn_type=ssl.SSLError) - await verify_fails('verify-full', - exn_type=ssl.SSLError) + await self._test_sslmode_fails('disable') + await self._test_sslmode_works('allow', expected_ssl=True) + await self._test_sslmode_works('prefer', expected_ssl=True) + await self._test_sslmode_fails('require', exc_type=ssl.SSLError) + await self._test_sslmode_fails('verify-ca', exc_type=ssl.SSLError) + await self._test_sslmode_fails( + 'verify-ca', host='127.0.0.1', exc_type=ssl.SSLError) + await self._test_sslmode_fails( + 'verify-full', exc_type=ssl.SSLError) async def test_sslmode_preserves_password_error(self): await self.con.execute( @@ -1951,12 +2410,7 @@ async def test_sslmode_preserves_password_error(self): for sslmode, first_type, fallback_type, fallback_is_ssl in cases: with self.subTest(sslmode=sslmode): self.cluster.reset_hba() - for address in ('127.0.0.0/24', '::1/128'): - self.cluster.add_hba_entry( - type=first_type, - address=ipaddress.ip_network(address), - database='postgres', user='ssl_user', - auth_method='password') + self._add_hba_entries(first_type, 'password') self.cluster.reload() connect_args = dict( @@ -1967,27 +2421,63 @@ async def test_sslmode_preserves_password_error(self): ssl=sslmode, ) - with self.assertRaisesRegex( - asyncpg.InvalidPasswordError, - 'password authentication failed', + async with self.connection_proxy() as (port, seen): + with self.assertRaisesRegex( + asyncpg.InvalidPasswordError, + 'password authentication failed', + ): + await self.connect(**connect_args, port=port) + self.assertEqual(seen, [not fallback_is_ssl, + fallback_is_ssl]) + + for fault, expected in ( + ({'disconnect': True}, asyncpg.InvalidPasswordError), + ({'error': asyncpg.InvalidCatalogNameError}, + asyncpg.InvalidCatalogNameError), ): - await self.connect(**connect_args) + async with self.connection_proxy({}, fault) as result: + port, seen = result + with self.assertRaises(expected) as raised: + await self.connect(**connect_args, port=port) + if expected is asyncpg.InvalidCatalogNameError: + self.assertIsInstance(raised.exception.__cause__, + asyncpg.InvalidPasswordError) + self.assertEqual(seen, [not fallback_is_ssl, + fallback_is_ssl]) # A password failure in the preferred mode must not prevent # a valid, differently-authenticated fallback connection. - for address in ('127.0.0.0/24', '::1/128'): - self.cluster.add_hba_entry( - type=fallback_type, - address=ipaddress.ip_network(address), - database='postgres', user='ssl_user', - auth_method='trust') + self._add_hba_entries(fallback_type, 'trust') self.cluster.reload() - con = await self.connect(**connect_args) - try: - self.assertEqual(con._protocol.is_ssl, fallback_is_ssl) - finally: - await con.close() + async with self.connection_proxy() as (port, seen): + await self._test_works( + **connect_args, port=port, + expected_ssl=fallback_is_ssl) + self.assertEqual(seen, [not fallback_is_ssl, + fallback_is_ssl]) + + async def test_ssl_client_authentication_error_is_terminal(self): + if self.cluster.get_pg_version() < (10, 0): + await self.con.execute("SET password_encryption = on") + else: + await self.con.execute("SET password_encryption = 'md5'") + await self.con.execute("ALTER ROLE ssl_user PASSWORD 'password'") + self.cluster.reset_hba() + self._add_hba_entries('host', 'md5', addresses=('127.0.0.0/24',)) + self.cluster.reload() + for mode in ('allow', 'prefer'): + async with self.connection_proxy() as (port, seen): + with unittest.mock.patch( + 'hashlib.md5', side_effect=ValueError('no md5') + ): + with self.assertRaisesRegex( + asyncpg.InternalClientError, 'no md5' + ): + await self.connect( + host='127.0.0.1', port=port, user='ssl_user', + password='password', ssl=mode) + self.assertEqual(seen, [mode == 'prefer']) async def test_ssl_connection_default_context(self): # XXX: uvloop artifact @@ -2006,26 +2496,9 @@ async def test_ssl_connection_pool(self): ssl_context = ssl.SSLContext(ssl.PROTOCOL_TLS_CLIENT) ssl_context.load_verify_locations(SSL_CA_CERT_FILE) - pool = await self.create_pool( - host='localhost', - user='ssl_user', - database='postgres', - min_size=5, - max_size=10, - ssl=ssl_context) - - async def worker(): - async with pool.acquire() as con: - self.assertEqual(await con.fetchval('SELECT 42'), 42) - - with self.assertRaises(asyncio.TimeoutError): - await con.execute('SELECT pg_sleep(5)', timeout=0.5) - - self.assertEqual(await con.fetchval('SELECT 43'), 43) - - tasks = [worker() for _ in range(100)] - await asyncio.gather(*tasks) - await pool.close() + await self._test_pool( + host='localhost', user='ssl_user', database='postgres', + ssl=ssl_context, expected_ssl=True) async def test_executemany_uvloop_ssl_issue_700(self): ssl_context = ssl.SSLContext(ssl.PROTOCOL_TLS_CLIENT) @@ -2089,16 +2562,12 @@ async def test_tls_version(self): '&ssl_min_protocol_version=TLSv1.2' '&ssl_max_protocol_version=TLSv1.1' ) - con = await self.connect( + await self._test_works( dsn='postgresql://ssl_user@localhost/postgres' '?sslmode=require' '&ssl_min_protocol_version=TLSv1.2' '&ssl_max_protocol_version=TLSv1.2' ) - try: - self.assertEqual(await con.fetchval('SELECT 42'), 42) - finally: - await con.close() finally: self.loop.set_exception_handler(old_handler) @@ -2106,15 +2575,7 @@ async def test_tls_version(self): @unittest.skipIf(os.environ.get('PGHOST'), 'unmanaged cluster') class TestClientSSLConnection(BaseTestSSLConnection): def _add_hba_entry(self): - self.cluster.add_hba_entry( - type='hostssl', address=ipaddress.ip_network('127.0.0.0/24'), - database='postgres', user='ssl_user', - auth_method='cert') - - self.cluster.add_hba_entry( - type='hostssl', address=ipaddress.ip_network('::1/128'), - database='postgres', user='ssl_user', - auth_method='cert') + self._add_hba_entries('hostssl', 'cert') async def test_ssl_connection_client_auth_fails_with_wrong_setup(self): ssl_context = ssl.create_default_context( @@ -2132,14 +2593,6 @@ async def test_ssl_connection_client_auth_fails_with_wrong_setup(self): ssl=ssl_context, ) - async def _test_works(self, **conn_args): - con = await self.connect(**conn_args) - - try: - self.assertEqual(await con.fetchval('SELECT 42'), 42) - finally: - await con.close() - async def test_ssl_connection_client_auth_custom_context(self): for key_file in (CLIENT_SSL_KEY_FILE, CLIENT_SSL_PROTECTED_KEY_FILE): ssl_context = ssl.create_default_context( @@ -2199,57 +2652,27 @@ async def test_ssl_connection_client_auth_dot_postgresql(self): @unittest.skipIf(os.environ.get('PGHOST'), 'unmanaged cluster') class TestNoSSLConnection(BaseTestSSLConnection): def _add_hba_entry(self): - self.cluster.add_hba_entry( - type='hostnossl', address=ipaddress.ip_network('127.0.0.0/24'), - database='postgres', user='ssl_user', - auth_method='trust') + self._add_hba_entries('hostnossl', 'trust') - self.cluster.add_hba_entry( - type='hostnossl', address=ipaddress.ip_network('::1/128'), - database='postgres', user='ssl_user', - auth_method='trust') + async def test_nossl_handshake_fallback(self): + if self.cluster.get_pg_version() < (12, 0): + self.skipTest('PostgreSQL < 12 cannot set SSL protocol version') + async with self.connection_proxy() as (port, seen): + await self._test_works( + dsn='postgresql://ssl_user@127.0.0.1/postgres' + '?sslmode=prefer&ssl_min_protocol_version=TLSv1.3', + port=port, expected_ssl=False) + self.assertEqual(seen, [True, False]) async def test_nossl_connection_sslmode(self): - async def verify_works(sslmode, *, host='localhost'): - con = None - try: - con = await self.connect( - dsn='postgresql://foo/postgres?sslmode=' + sslmode, - host=host, - user='ssl_user') - self.assertEqual(await con.fetchval('SELECT 42'), 42) - self.assertFalse(con._protocol.is_ssl) - finally: - if con: - await con.close() - - async def verify_fails(sslmode, *, host='localhost'): - # XXX: uvloop artifact - old_handler = self.loop.get_exception_handler() - con = None - try: - self.loop.set_exception_handler(lambda *args: None) - with self.assertRaises( - asyncpg.InvalidAuthorizationSpecificationError - ): - con = await self.connect( - dsn='postgresql://foo/?sslmode=' + sslmode, - host=host, - user='ssl_user') - await con.fetchval('SELECT 42') - finally: - if con: - await con.close() - self.loop.set_exception_handler(old_handler) - - await verify_works('disable') - await verify_works('allow') - await verify_works('prefer') - await verify_fails('require') + await self._test_sslmode_works('disable', expected_ssl=False) + await self._test_sslmode_works('allow', expected_ssl=False) + await self._test_sslmode_works('prefer', expected_ssl=False) + await self._test_sslmode_fails('require') with mock_dot_postgresql(): - await verify_fails('require') - await verify_fails('verify-ca') - await verify_fails('verify-full') + await self._test_sslmode_fails('require') + await self._test_sslmode_fails('verify-ca') + await self._test_sslmode_fails('verify-full') async def test_nossl_connection_prefer_cancel(self): con = await self.connect( @@ -2258,35 +2681,14 @@ async def test_nossl_connection_prefer_cancel(self): user='ssl_user') try: self.assertFalse(con._protocol.is_ssl) - with self.assertRaises(asyncio.TimeoutError): - await con.execute('SELECT pg_sleep(5)', timeout=0.5) - val = await con.fetchval('SELECT 123') - self.assertEqual(val, 123) + await self._test_cancellation_recovery(con) finally: await con.close() async def test_nossl_connection_pool(self): - pool = await self.create_pool( - host='localhost', - user='ssl_user', - database='postgres', - min_size=5, - max_size=10, - ssl='prefer') - - async def worker(): - async with pool.acquire() as con: - self.assertFalse(con._protocol.is_ssl) - self.assertEqual(await con.fetchval('SELECT 42'), 42) - - with self.assertRaises(asyncio.TimeoutError): - await con.execute('SELECT pg_sleep(5)', timeout=0.5) - - self.assertEqual(await con.fetchval('SELECT 43'), 43) - - tasks = [worker() for _ in range(100)] - await asyncio.gather(*tasks) - await pool.close() + await self._test_pool( + host='localhost', user='ssl_user', database='postgres', + ssl='prefer', expected_ssl=False) class TestConnectionGC(tb.ClusterTestCase):