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