From 2a4c369d58d7c066a0fcdad01783656c3618662e Mon Sep 17 00:00:00 2001 From: Brian Martin Date: Tue, 4 Jun 2019 22:20:11 -0400 Subject: [PATCH] add more xmpp tests - test xmpp server - test starttls upgrade --- bumper/xmppserver.py | 7 +-- tests/test_mqttserver.py | 3 +- tests/test_xmppserver.py | 103 ++++++++++++++++++++++++++++++++++++++- 3 files changed, 107 insertions(+), 6 deletions(-) diff --git a/bumper/xmppserver.py b/bumper/xmppserver.py index 0dabdc9..453de22 100644 --- a/bumper/xmppserver.py +++ b/bumper/xmppserver.py @@ -14,6 +14,7 @@ class XMPPServer: server_id = "ecouser.net" clients = [] exit_flag = False + server = None def __init__(self, address): # Initialize bot server @@ -27,12 +28,11 @@ class XMPPServer: loop = asyncio.get_running_loop() - server = await loop.create_server( + self.server = await loop.create_server( self.xmpp_protocol, host=self.address[0], port=self.address[1] ) - async with server: - await server.serve_forever() + self.server_coro = loop.create_task(self.server.serve_forever()) def disconnect(self): try: @@ -42,6 +42,7 @@ class XMPPServer: self.exit_flag = True xmppserverlog.debug("shutting down") + self.server_coro.cancel() except Exception as e: xmppserverlog.error("{}".format(e)) diff --git a/tests/test_mqttserver.py b/tests/test_mqttserver.py index 848ed09..24ef16f 100644 --- a/tests/test_mqttserver.py +++ b/tests/test_mqttserver.py @@ -365,8 +365,7 @@ async def test_mqttserver(): fake_bot.Client._connected_state._value == True ) # Check fake_bot is connected await fake_bot.Client.disconnect() - - await mqtt_helperbot.Client.reconnect() # This forces the above disconnect + await asyncio.sleep(0.1) await mqtt_server.broker.shutdown() \ No newline at end of file diff --git a/tests/test_xmppserver.py b/tests/test_xmppserver.py index e56fdcf..6ed558a 100644 --- a/tests/test_xmppserver.py +++ b/tests/test_xmppserver.py @@ -7,6 +7,9 @@ import json import tinydb import pytest_asyncio import xml.etree.ElementTree as ET +import socket +from testfixtures import LogCapture +import ssl def return_send_data(data, *args, **kwargs): @@ -14,7 +17,45 @@ def return_send_data(data, *args, **kwargs): def mock_transport_extra_info(*args, **kwargs): - return ("127.0.0.1", 1234) + return ("127.0.0.1", 5223) + + +async def test_xmpp_server(): + with LogCapture("xmppserver") as l: + xmpp_address = ("127.0.0.1", 5223) + xmpp_server = bumper.XMPPServer(xmpp_address) + await xmpp_server.start_async_server() + + reader, writer = await asyncio.open_connection("127.0.0.1", 5223) + + writer.write(b"") # Start stream + await writer.drain() + + await asyncio.sleep(0.1) + + assert len(xmpp_server.clients) == 1 # Client count increased + assert ( + xmpp_server.clients[0].address[1] + == writer.transport.get_extra_info("sockname")[1] + ) + + writer.close() # Close connection + await writer.wait_closed() + + await asyncio.sleep(0.1) + + assert len(xmpp_server.clients) == 0 # Client count decreased + + reader, writer = await asyncio.open_connection("127.0.0.1", 5223) + + writer.write(b"") # Start stream + await writer.drain() + + await asyncio.sleep(0.1) + xmpp_server.disconnect() + await asyncio.sleep(0.1) + assert len(xmpp_server.clients) == 0 # Client count decreased + print(l) async def test_client_connect_no_starttls(*args, **kwargs): @@ -167,6 +208,66 @@ async def test_client_connect_starttls_called(*args, **kwargs): assert xmppclient.state == xmppclient.INIT # Client moved to INIT state +async def test_xmpp_server_client_tls(): + with LogCapture("xmppserver") as l: + + async def do_stuff_after_start_tls( + ssl_reader, ssl_writer + ): # Used after starttls + + writer.write( + b"" + ) # New Stream + + await writer.drain() + + await asyncio.sleep(0.1) + + writer.write( + b'AGZ1aWRfdG1wdXNlcgAwL0lPU0Y1M0QwN0JBL3VzXzg5ODgwMmZkYmM0NDQxYjBiYzgxNWIxZDFjNjgzMDJl' + ) # Send Auth + + await writer.drain() + + await asyncio.sleep(0.1) + + xmpp_address = ("127.0.0.1", 5223) + xmpp_server = bumper.XMPPServer(xmpp_address) + await xmpp_server.start_async_server() + + reader, writer = await asyncio.open_connection("127.0.0.1", 5223) + + writer.write( + b"" + ) # Start stream + await writer.drain() + + await asyncio.sleep(0.1) + + writer.write( + b"" + ) # Send StartTLS + await writer.drain() + + await asyncio.sleep(0.1) + + # Below will upgrade connection to TLS then callback to "do_stuff_after_start_tls" + ssl_context = ssl.create_default_context(ssl.Purpose.SERVER_AUTH) + ssl_context.check_hostname = False + ssl_context.load_verify_locations(cafile=bumper.ca_cert) + loop = asyncio.get_event_loop() + transport = writer.transport + protocol = writer.transport.get_protocol() + new_transport = await loop.start_tls( + transport, protocol, ssl_context, server_side=False + ) + protocol._stream_reader = asyncio.StreamReader(loop=loop) + protocol._client_connected_cb = do_stuff_after_start_tls + protocol.connection_made(new_transport) + + print(l) + + async def test_client_init(*args, **kwargs): test_transport = asyncio.Transport() test_transport.get_extra_info = mock.Mock(return_value=mock_transport_extra_info())