Close client connections if main thread is stopped
This commit is contained in:
parent
8a59115cad
commit
878310d3e1
1 changed files with 44 additions and 7 deletions
|
|
@ -29,6 +29,7 @@ import select
|
||||||
import struct
|
import struct
|
||||||
import sys
|
import sys
|
||||||
import threading
|
import threading
|
||||||
|
import time
|
||||||
import logging
|
import logging
|
||||||
from collections import namedtuple
|
from collections import namedtuple
|
||||||
|
|
||||||
|
|
@ -60,9 +61,7 @@ AUTH_METHODS_REVERSE = {val: key for key, val in AUTH_METHODS.items()}
|
||||||
|
|
||||||
def main():
|
def main():
|
||||||
args = get_args()
|
args = get_args()
|
||||||
|
|
||||||
configure_logger(args.logging_level)
|
configure_logger(args.logging_level)
|
||||||
|
|
||||||
target_backend = args.backend
|
target_backend = args.backend
|
||||||
|
|
||||||
sock = socket.socket(socket.AF_INET, socket.SOCK_STREAM)
|
sock = socket.socket(socket.AF_INET, socket.SOCK_STREAM)
|
||||||
|
|
@ -72,18 +71,29 @@ def main():
|
||||||
backlog = 5
|
backlog = 5
|
||||||
|
|
||||||
sock.listen(backlog)
|
sock.listen(backlog)
|
||||||
|
|
||||||
_logger.info('Listening for connections')
|
_logger.info('Listening for connections')
|
||||||
|
|
||||||
|
# Maintain a reference to running threads to enable stopping them when the
|
||||||
|
# script terminates
|
||||||
|
threads = set()
|
||||||
|
|
||||||
|
last_check_for_stopped_threads = time.time()
|
||||||
|
|
||||||
try:
|
try:
|
||||||
while True:
|
while True:
|
||||||
client_socket, address = sock.accept()
|
client_socket, address = sock.accept()
|
||||||
client_handler = ClientConnection(client_socket, target_backend)
|
client_handler = ClientConnection(client_socket, target_backend)
|
||||||
client_handler.start()
|
client_handler.start()
|
||||||
|
|
||||||
|
threads.add(client_handler)
|
||||||
|
if time.time() - last_check_for_stopped_threads > 3:
|
||||||
|
remove_stopped_threads(threads)
|
||||||
|
|
||||||
except (KeyboardInterrupt, SystemExit):
|
except (KeyboardInterrupt, SystemExit):
|
||||||
_logger.info('Received exit, shutting down sockets')
|
_logger.info('Received exit, shutting down sockets')
|
||||||
sock.shutdown(socket.SHUT_RDWR)
|
sock.shutdown(socket.SHUT_RDWR)
|
||||||
sock.close()
|
sock.close()
|
||||||
|
stop_threads(threads)
|
||||||
|
|
||||||
|
|
||||||
def get_args():
|
def get_args():
|
||||||
|
|
@ -99,6 +109,20 @@ def configure_logger(level):
|
||||||
_logger.setLevel(getattr(logging, level.upper()))
|
_logger.setLevel(getattr(logging, level.upper()))
|
||||||
|
|
||||||
|
|
||||||
|
def remove_stopped_threads(threads):
|
||||||
|
threads_to_remove = []
|
||||||
|
for thread in threads:
|
||||||
|
if thread.stopped:
|
||||||
|
threads_to_remove.append(thread)
|
||||||
|
for thread in threads_to_remove:
|
||||||
|
threads.remove(thread)
|
||||||
|
|
||||||
|
|
||||||
|
def stop_threads(threads):
|
||||||
|
for thread in threads:
|
||||||
|
thread.stop()
|
||||||
|
|
||||||
|
|
||||||
class ClientConnection(threading.Thread):
|
class ClientConnection(threading.Thread):
|
||||||
|
|
||||||
def __init__(self, client_socket, target_backend):
|
def __init__(self, client_socket, target_backend):
|
||||||
|
|
@ -109,8 +133,18 @@ class ClientConnection(threading.Thread):
|
||||||
self.packet_handler = self.handle_startup_packet
|
self.packet_handler = self.handle_startup_packet
|
||||||
self.target_backend = target_backend
|
self.target_backend = target_backend
|
||||||
self.server_socket = None
|
self.server_socket = None
|
||||||
|
self._stop = threading.Event()
|
||||||
|
|
||||||
|
|
||||||
|
def stop(self):
|
||||||
|
self._stop.set()
|
||||||
|
|
||||||
|
|
||||||
|
@property
|
||||||
|
def stopped(self):
|
||||||
|
return self._stop.isSet()
|
||||||
|
|
||||||
|
|
||||||
|
|
||||||
def run(self):
|
def run(self):
|
||||||
_logger.debug('Thread running')
|
_logger.debug('Thread running')
|
||||||
# wait for startup packet
|
# wait for startup packet
|
||||||
|
|
@ -124,7 +158,7 @@ class ClientConnection(threading.Thread):
|
||||||
|
|
||||||
data = self.socket.read()
|
data = self.socket.read()
|
||||||
try:
|
try:
|
||||||
while data:
|
while data and not self.stopped:
|
||||||
if not self.packet_handler(data):
|
if not self.packet_handler(data):
|
||||||
break
|
break
|
||||||
if self.server_socket:
|
if self.server_socket:
|
||||||
|
|
@ -135,8 +169,10 @@ class ClientConnection(threading.Thread):
|
||||||
_logger.warning('No more data received')
|
_logger.warning('No more data received')
|
||||||
raise Exception('No more data')
|
raise Exception('No more data')
|
||||||
_logger.debug('Using select to wait for data')
|
_logger.debug('Using select to wait for data')
|
||||||
while True:
|
while not self.stopped:
|
||||||
sockets_with_data = select.select([self.server_socket, self.socket], [], [])[0]
|
timeout = 0
|
||||||
|
sockets = [self.server_socket, self.socket]
|
||||||
|
sockets_with_data = select.select(sockets, [], [], timeout)[0]
|
||||||
if self.socket in sockets_with_data:
|
if self.socket in sockets_with_data:
|
||||||
data = self.socket.read()
|
data = self.socket.read()
|
||||||
_logger.debug('Client -> Server: %s' % repr(data))
|
_logger.debug('Client -> Server: %s' % repr(data))
|
||||||
|
|
@ -264,6 +300,7 @@ class ClientConnection(threading.Thread):
|
||||||
self.socket.close()
|
self.socket.close()
|
||||||
except:
|
except:
|
||||||
_logger.exception('Got exception when trying to close socket')
|
_logger.exception('Got exception when trying to close socket')
|
||||||
|
self.stop()
|
||||||
|
|
||||||
|
|
||||||
def socket_is_closed(sock):
|
def socket_is_closed(sock):
|
||||||
|
|
|
||||||
Loading…
Add table
Reference in a new issue