Add unit tests #39

Merged
bmartin5692 merged 9 commits from add-tests into master 2019-06-05 05:51:23 +02:00
3 changed files with 107 additions and 6 deletions
Showing only changes of commit 2a4c369d58 - Show all commits

View file

@ -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))

View file

@ -366,7 +366,6 @@ async def test_mqttserver():
) # 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()

View file

@ -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"<stream:stream />") # 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"<stream:stream />") # 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"<stream:stream xmlns='jabber:client' xmlns:stream='http://etherx.jabber.org/streams' version='1.0' to='ecouser.net'>"
) # New Stream
await writer.drain()
await asyncio.sleep(0.1)
writer.write(
b'<auth xmlns="urn:ietf:params:xml:ns:xmpp-sasl" mechanism="PLAIN">AGZ1aWRfdG1wdXNlcgAwL0lPU0Y1M0QwN0JBL3VzXzg5ODgwMmZkYmM0NDQxYjBiYzgxNWIxZDFjNjgzMDJl</auth>'
) # 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"<stream:stream xmlns='jabber:client' xmlns:stream='http://etherx.jabber.org/streams' version='1.0' to='ecouser.net'>"
) # Start stream
await writer.drain()
await asyncio.sleep(0.1)
writer.write(
b"<starttls xmlns='urn:ietf:params:xml:ns:xmpp-tls'/>"
) # 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())