diff --git a/asyncpg/connect_utils.py b/asyncpg/connect_utils.py index 07c4fdde..70b7a6ba 100644 --- a/asyncpg/connect_utils.py +++ b/asyncpg/connect_utils.py @@ -934,6 +934,12 @@ 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'): + message = data[1:].rstrip(b'\x00\r\n').decode( + 'utf-8', errors='replace') + self.on_data.set_exception( + exceptions.InterfaceError( + message or 'server error during SSL negotiation')) else: self.on_data.set_exception( ConnectionError( diff --git a/tests/test_connect.py b/tests/test_connect.py index 955fb825..741d89f8 100644 --- a/tests/test_connect.py +++ b/tests/test_connect.py @@ -99,6 +99,19 @@ def mock_dev_null_home_dir(): yield +class TestTLSUpgradeProto(tb.TestCase): + + async def test_error_response_preserves_server_message(self): + context = ssl.SSLContext(ssl.PROTOCOL_TLS_CLIENT) + proto = connect_utils.TLSUpgradeProto( + self.loop, 'localhost', 5432, context, False) + + proto.data_received(b'Etoo many connections\n\x00') + with self.assertRaisesRegex( + exceptions.InterfaceError, 'too many connections'): + await proto.on_data + + class TestSettings(tb.ConnectedTestCase): async def test_get_settings_01(self):