Added test to check that child threads die when parent dies

This commit is contained in:
Roy Hyunjin Han 2013-02-09 19:12:21 -08:00
commit 0a5b069cdd
7 changed files with 193 additions and 137 deletions

View file

@ -1,3 +1,6 @@
0.4
---
0.3 0.3
--- ---
- Added support for secure connections - Added support for secure connections

View file

@ -35,10 +35,9 @@ Activate isolated environment. ::
Emit. :: Emit. ::
from socketIO_client import SocketIO from socketIO_client import SocketIO
with SocketIO('localhost', 8000) as socketIO:
socketIO = SocketIO('localhost', 8000) socketIO.emit('aaa')
socketIO.emit('aaa', {'bbb': 'ccc'}) socketIO.wait(1) # Wait a second
socketIO.wait(seconds=1) # Exit after one second
Emit with callback. :: Emit with callback. ::
@ -47,9 +46,9 @@ Emit with callback. ::
def on_response(*args): def on_response(*args):
print args print args
socketIO = SocketIO('localhost', 8000) with SocketIO('localhost', 8000) as socketIO:
socketIO.emit('aaa', {'bbb': 'ccc'}, on_response) socketIO.emit('aaa', {'bbb': 'ccc'}, on_response)
socketIO.wait(forCallbacks=True) # Exit after callbacks run socketIO.wait(seconds=1, forCallbacks=True) # Wait for callback
Define events. :: Define events. ::

View file

@ -1,5 +1,8 @@
Let user define a proxy #5 + Fix unittests
Let user emit without arguments #5 + Fix exceptions when websocket server disappears
Fix thread exceptions
Integrate Zac's fork #6 Integrate Zac's fork #6
Integrate Sajal's fork #7 Integrate Sajal's fork #7
Integrate Francis's fork #10 Integrate Francis's fork #10

17
serve_tests.py Executable file → Normal file
View file

@ -1,7 +1,14 @@
'Launch this server in another terminal window before running tests' 'Launch this server in another terminal window before running tests'
from socketio import socketio_manage import sys
from socketio.namespace import BaseNamespace try:
from socketio.server import SocketIOServer from socketio import socketio_manage
from socketio.namespace import BaseNamespace
from socketio.server import SocketIOServer
except ImportError:
from setuptools.command import easy_install
easy_install.main(['-U', 'gevent-socketio'])
print('\nPlease run the script again to launch the test server.')
sys.exit(1)
class Namespace(BaseNamespace): class Namespace(BaseNamespace):
@ -25,5 +32,7 @@ class Application(object):
if __name__ == '__main__': if __name__ == '__main__':
socketIOServer = SocketIOServer(('0.0.0.0', 8000), Application()) port = 8000
print 'Starting server at port %s' % port
socketIOServer = SocketIOServer(('0.0.0.0', port), Application())
socketIOServer.serve_forever() socketIOServer.serve_forever()

1
setup.py Executable file → Normal file
View file

@ -24,7 +24,6 @@ setup(
url='https://github.com/invisibleroads/socketIO-client', url='https://github.com/invisibleroads/socketIO-client',
install_requires=[ install_requires=[
'anyjson', 'anyjson',
'gevent-socketio',
'websocket-client', 'websocket-client',
], ],
packages=find_packages(), packages=find_packages(),

View file

@ -1,11 +1,16 @@
import websocket import sys
import traceback
import socket
from anyjson import dumps, loads from anyjson import dumps, loads
from functools import partial
from threading import Thread, Event from threading import Thread, Event
from time import sleep from time import sleep
from urllib import urlopen from urllib import urlopen
from websocket import WebSocketConnectionClosedException, create_connection
__version__ = '0.3' __version__ = '0.4'
PROTOCOL = 1 # SocketIO protocol version PROTOCOL = 1 # SocketIO protocol version
@ -16,7 +21,7 @@ class BaseNamespace(object): # pragma: no cover
def __init__(self, socketIO): def __init__(self, socketIO):
self.socketIO = socketIO self.socketIO = socketIO
def on_connect(self, socketIO): def on_connect(self):
pass pass
def on_disconnect(self): def on_disconnect(self):
@ -46,54 +51,69 @@ class BaseNamespace(object): # pragma: no cover
class SocketIO(object): class SocketIO(object):
messageID = 0 _messageID = 0
def __init__(self, host, port, Namespace=BaseNamespace, secure=False): def __init__(self, host, port, Namespace=BaseNamespace, secure=False, proxies=None):
self.host = host self._host = host
self.port = int(port) self._port = int(port)
self.namespace = Namespace(self) self._namespace = Namespace(self)
self.secure = secure self._secure = secure
self.__connect() self._proxies = proxies
self._connect()
heartbeatInterval = self.heartbeatTimeout - 2 heartbeatInterval = self._heartbeatTimeout - 2
self.heartbeatThread = RhythmicThread(heartbeatInterval, self._heartbeatThread = RhythmicThread(heartbeatInterval, self._send_heartbeat)
self._send_heartbeat) self._heartbeatThread.start()
self.heartbeatThread.start()
self.channelByName = {} self._channelByName = {}
self.callbackByEvent = {} self._callbackByEvent = {}
self.namespaceThread = ListenerThread(self) self._namespaceThread = ListenerThread(self._recv_packet, self._get_callback)
self.namespaceThread.start() self._namespaceThread.start()
def __del__(self): # pragma: no cover def __enter__(self):
self.heartbeatThread.cancel() return self
self.namespaceThread.cancel()
self.connection.close()
def __connect(self): def __exit__(self, exc_type, exc_value, traceback):
baseURL = '%s:%d/socket.io/%s' % (self.host, self.port, PROTOCOL) self.__del__()
def __del__(self):
self._heartbeatThread.cancel()
self._namespaceThread.cancel()
self._connection.close()
def _connect(self):
baseURL = '%s:%d/socket.io/%s' % (self._host, self._port, PROTOCOL)
try: try:
response = urlopen('%s://%s/' % ( response = urlopen('%s://%s/' % (
'https' if self.secure else 'http', baseURL)) 'https' if self._secure else 'http', baseURL),
proxies=self._proxies)
except IOError: # pragma: no cover except IOError: # pragma: no cover
raise SocketIOError('Could not start connection') raise SocketIOError('Could not start connection')
if 200 != response.getcode(): # pragma: no cover if 200 != response.getcode(): # pragma: no cover
raise SocketIOError('Could not establish connection') raise SocketIOError('Could not establish connection')
responseParts = response.readline().split(':') responseParts = response.readline().split(':')
self.sessionID = responseParts[0] self._sessionID = responseParts[0]
self.heartbeatTimeout = int(responseParts[1]) self._heartbeatTimeout = int(responseParts[1])
self.connectionTimeout = int(responseParts[2]) self._connectionTimeout = int(responseParts[2])
self.supportedTransports = responseParts[3].split(',') self._supportedTransports = responseParts[3].split(',')
if 'websocket' not in self.supportedTransports: if 'websocket' not in self._supportedTransports:
raise SocketIOError('Could not parse handshake') # pragma: no cover raise SocketIOError('Could not parse handshake') # pragma: no cover
socketURL = '%s://%s/websocket/%s' % ( socketURL = '%s://%s/websocket/%s' % (
'wss' if self.secure else 'ws', baseURL, self.sessionID) 'wss' if self._secure else 'ws', baseURL, self._sessionID)
self.connection = websocket.create_connection(socketURL) self._connection = create_connection(socketURL)
def _recv_packet(self): def _recv_packet(self):
code, packetID, channelName, data = -1, None, None, None code, packetID, channelName, data = -1, None, None, None
packet = self.connection.recv() try:
packet = self._connection.recv()
except WebSocketConnectionClosedException:
raise SocketIOConnectionError('Lost connection (Connection closed)')
except socket.timeout:
raise SocketIOConnectionError('Lost connection (Connection timed out)')
try:
packetParts = packet.split(':', 3) packetParts = packet.split(':', 3)
except AttributeError:
raise SocketIOPacketError('Received invalid packet (%s)' % packet)
packetCount = len(packetParts) packetCount = len(packetParts)
if 4 == packetCount: if 4 == packetCount:
code, packetID, channelName, data = packetParts code, packetID, channelName, data = packetParts
@ -104,34 +124,36 @@ class SocketIO(object):
return int(code), packetID, channelName, data return int(code), packetID, channelName, data
def _send_packet(self, code, channelName='', data='', callback=None): def _send_packet(self, code, channelName='', data='', callback=None):
self.connection.send(':'.join([ callbackNumber = self._set_callback(callback) if callback else ''
str(code), packetParts = [str(code), callbackNumber, channelName, data]
self.set_callback(callback) if callback else '', try:
channelName, self._connection.send(':'.join(packetParts))
data])) except socket.error:
raise SocketIOPacketError('Could not send packet')
def disconnect(self, channelName=''): def disconnect(self, channelName=''):
self._send_packet(0, channelName) self._send_packet(0, channelName)
if channelName: if channelName:
del self.channelByName[channelName] del self._channelByName[channelName]
else: else:
self.__del__() self.__del__()
@property @property
def connected(self): def connected(self):
return self.connection.connected return self._connection.connected
def connect(self, channelName, Namespace=BaseNamespace): def connect(self, channelName, Namespace=BaseNamespace):
channel = Channel(self, channelName, Namespace) channel = Channel(self, channelName, Namespace)
self.channelByName[channelName] = channel self._channelByName[channelName] = channel
self._send_packet(1, channelName) self._send_packet(1, channelName)
return channel return channel
def _send_heartbeat(self): def _send_heartbeat(self):
try: try:
self._send_packet(2) self._send_packet(2)
except: except SocketIOPacketError:
self.__del__() print 'Could not send heartbeat'
pass
def message(self, messageData, callback=None, channelName=''): def message(self, messageData, callback=None, channelName=''):
if isinstance(messageData, basestring): if isinstance(messageData, basestring):
@ -144,40 +166,39 @@ class SocketIO(object):
def emit(self, eventName, *eventArguments, **eventKeywords): def emit(self, eventName, *eventArguments, **eventKeywords):
code = 5 code = 5
if callable(eventArguments[-1]): callback = None
if eventArguments and callable(eventArguments[-1]):
callback = eventArguments[-1] callback = eventArguments[-1]
eventArguments = eventArguments[:-1] eventArguments = eventArguments[:-1]
else:
callback = None
channelName = eventKeywords.get('channelName', '') channelName = eventKeywords.get('channelName', '')
data = dumps(dict(name=eventName, args=eventArguments)) data = dumps(dict(name=eventName, args=eventArguments))
self._send_packet(code, channelName, data, callback) self._send_packet(code, channelName, data, callback)
def get_callback(self, channelName, eventName): def _get_callback(self, channelName, eventName):
'Get callback associated with channelName and eventName' 'Get callback associated with channelName and eventName'
socketIO = self.channelByName[channelName] if channelName else self socketIO = self._channelByName[channelName] if channelName else self
try: try:
return socketIO.callbackByEvent[eventName] return socketIO._callbackByEvent[eventName]
except KeyError: except KeyError:
pass pass
namespace = socketIO.namespace
def callback_(*eventArguments): def callback_(*eventArguments):
return namespace.on_(eventName, *eventArguments) return socketIO._namespace.on_(eventName, *eventArguments)
return getattr(namespace, name_callback(eventName), callback_) callbackName = 'on_' + eventName.replace(' ', '_')
return getattr(socketIO._namespace, callbackName, callback_)
def set_callback(self, callback): def _set_callback(self, callback):
'Set callback that will be called after receiving an acknowledgment' 'Set callback that will be called after receiving an acknowledgment'
self.messageID += 1 self._messageID += 1
self.namespaceThread.set_callback(self.messageID, callback) self._namespaceThread.set_callback(self._messageID, callback)
return '%s+' % self.messageID return '%s+' % self._messageID
def on(self, eventName, callback): def on(self, eventName, callback):
self.callbackByEvent[eventName] = callback self._callbackByEvent[eventName] = callback
def wait(self, seconds=None, forCallbacks=False): def wait(self, seconds=None, forCallbacks=False):
if forCallbacks: if forCallbacks:
self.namespaceThread.wait_for_callbacks(seconds) self._namespaceThread.wait_for_callbacks(seconds)
elif seconds: elif seconds:
sleep(seconds) sleep(seconds)
else: else:
@ -191,24 +212,22 @@ class SocketIO(object):
class Channel(object): class Channel(object):
def __init__(self, socketIO, channelName, Namespace): def __init__(self, socketIO, channelName, Namespace):
self.socketIO = socketIO self._socketIO = socketIO
self.channelName = channelName self._channelName = channelName
self.namespace = Namespace(self) self._namespace = Namespace(self)
self.callbackByEvent = {} self._callbackByEvent = {}
def disconnect(self): def disconnect(self):
self.socketIO.disconnect(self.channelName) self._socketIO.disconnect(self._channelName)
def emit(self, eventName, *eventArguments): def emit(self, eventName, *eventArguments):
self.socketIO.emit(eventName, *eventArguments, self._socketIO.emit(eventName, *eventArguments, channelName=self._channelName)
channelName=self.channelName)
def message(self, messageData, callback=None): def message(self, messageData, callback=None):
self.socketIO.message(messageData, callback, self._socketIO.message(messageData, callback, channelName=self._channelName)
channelName=self.channelName)
def on(self, eventName, eventCallback): def on(self, eventName, eventCallback):
self.callbackByEvent[eventName] = eventCallback self._callbackByEvent[eventName] = eventCallback
class ListenerThread(Thread): class ListenerThread(Thread):
@ -216,20 +235,26 @@ class ListenerThread(Thread):
daemon = True daemon = True
def __init__(self, socketIO): def __init__(self, recv_packet, get_callback):
super(ListenerThread, self).__init__() super(ListenerThread, self).__init__()
self.socketIO = socketIO
self.done = Event() self.done = Event()
self.waitingForCallbacks = Event() self.waitingForCallbacks = Event()
self.callbackByMessageID = {} self.callbackByMessageID = {}
self.get_callback = self.socketIO.get_callback self.recv_packet = recv_packet
self.get_callback = get_callback
def run(self): def run(self):
try:
while not self.done.is_set(): while not self.done.is_set():
try: try:
code, packetID, channelName, data = self.socketIO._recv_packet() code, packetID, channelName, data = self.recv_packet()
except: except SocketIOConnectionError, error:
print error
return
except SocketIOPacketError, error:
print error
continue continue
get_channel_callback = partial(self.get_callback, channelName)
try: try:
delegate = { delegate = {
0: self.on_disconnect, 0: self.on_disconnect,
@ -243,7 +268,10 @@ class ListenerThread(Thread):
}[code] }[code]
except KeyError: except KeyError:
continue continue
delegate(packetID, channelName, data) delegate(packetID, get_channel_callback, data)
except:
exc_type, exc_value, exc_traceback = sys.exc_info()
open('tracebacks.log', 'a+t').write('\n'.join(traceback.format_tb(exc_traceback)))
def cancel(self): def cancel(self):
self.done.set() self.done.set()
@ -255,33 +283,28 @@ class ListenerThread(Thread):
def set_callback(self, messageID, callback): def set_callback(self, messageID, callback):
self.callbackByMessageID[messageID] = callback self.callbackByMessageID[messageID] = callback
def on_disconnect(self, packetID, channelName, data): def on_disconnect(self, packetID, get_channel_callback, data):
callback = self.get_callback(channelName, 'disconnect') get_channel_callback('disconnect')()
callback()
def on_connect(self, packetID, channelName, data): def on_connect(self, packetID, get_channel_callback, data):
callback = self.get_callback(channelName, 'connect') get_channel_callback('connect')()
callback(self.socketIO)
def on_heartbeat(self, packetID, channelName, data): def on_heartbeat(self, packetID, get_channel_callback, data):
pass pass
def on_message(self, packetID, channelName, data): def on_message(self, packetID, get_channel_callback, data):
callback = self.get_callback(channelName, 'message') get_channel_callback('message')(data)
callback(data)
def on_json(self, packetID, channelName, data): def on_json(self, packetID, get_channel_callback, data):
callback = self.get_callback(channelName, 'message') get_channel_callback('message')(loads(data))
callback(loads(data))
def on_event(self, packetID, channelName, data): def on_event(self, packetID, get_channel_callback, data):
valueByName = loads(data) valueByName = loads(data)
eventName = valueByName['name'] eventName = valueByName['name']
eventArguments = valueByName['args'] eventArguments = valueByName['args']
callback = self.get_callback(channelName, eventName) get_channel_callback(eventName)(*eventArguments)
callback(*eventArguments)
def on_acknowledgment(self, packetID, channelName, data): def on_acknowledgment(self, packetID, get_channel_callback, data):
dataParts = data.split('+', 1) dataParts = data.split('+', 1)
messageID = int(dataParts[0]) messageID = int(dataParts[0])
arguments = loads(dataParts[1]) or [] arguments = loads(dataParts[1]) or []
@ -296,21 +319,20 @@ class ListenerThread(Thread):
if self.waitingForCallbacks.is_set() and not callbackCount: if self.waitingForCallbacks.is_set() and not callbackCount:
self.cancel() self.cancel()
def on_error(self, packetID, channelName, data): def on_error(self, packetID, get_channel_callback, data):
reason, advice = data.split('+', 1) reason, advice = data.split('+', 1)
callback = self.get_callback(channelName, 'error') get_channel_callback('error')(reason, advice)
callback(reason, advice)
class RhythmicThread(Thread): class RhythmicThread(Thread):
'Execute rhythmicFunction every few seconds' 'Execute call every few seconds'
daemon = True daemon = True
def __init__(self, intervalInSeconds, rhythmicFunction, *args, **kw): def __init__(self, intervalInSeconds, call, *args, **kw):
super(RhythmicThread, self).__init__() super(RhythmicThread, self).__init__()
self.intervalInSeconds = intervalInSeconds self.intervalInSeconds = intervalInSeconds
self.rhythmicFunction = rhythmicFunction self.call = call
self.args = args self.args = args
self.kw = kw self.kw = kw
self.done = Event() self.done = Event()
@ -318,10 +340,11 @@ class RhythmicThread(Thread):
def run(self): def run(self):
try: try:
while not self.done.is_set(): while not self.done.is_set():
self.rhythmicFunction(*self.args, **self.kw) self.call(*self.args, **self.kw)
self.done.wait(self.intervalInSeconds) self.done.wait(self.intervalInSeconds)
except: except:
pass exc_type, exc_value, exc_traceback = sys.exc_info()
open('tracebacks.log', 'a+t').write('\n'.join(traceback.format_tb(exc_traceback)))
def cancel(self): def cancel(self):
self.done.set() self.done.set()
@ -331,5 +354,9 @@ class SocketIOError(Exception):
pass pass
def name_callback(eventName): class SocketIOConnectionError(SocketIOError):
return 'on_' + eventName.replace(' ', '_') pass
class SocketIOPacketError(SocketIOError):
pass

View file

@ -15,10 +15,16 @@ class TestSocketIO(TestCase):
self.assertEqual(socketIO.connected, False) self.assertEqual(socketIO.connected, False)
def test_emit(self): def test_emit(self):
socketIO = SocketIO('localhost', 8000, Namespace)
socketIO.emit('aaa')
sleep(0.5)
self.assertEqual(socketIO._namespace.payload, '')
def test_emit_with_payload(self):
socketIO = SocketIO('localhost', 8000, Namespace) socketIO = SocketIO('localhost', 8000, Namespace)
socketIO.emit('aaa', PAYLOAD) socketIO.emit('aaa', PAYLOAD)
sleep(0.5) sleep(0.5)
self.assertEqual(socketIO.namespace.payload, PAYLOAD) self.assertEqual(socketIO._namespace.payload, PAYLOAD)
def test_emit_with_callback(self): def test_emit_with_callback(self):
global ON_RESPONSE_CALLED global ON_RESPONSE_CALLED
@ -43,16 +49,26 @@ class TestSocketIO(TestCase):
newsSocket = mainSocket.connect('/news', Namespace) newsSocket = mainSocket.connect('/news', Namespace)
newsSocket.emit('aaa', PAYLOAD) newsSocket.emit('aaa', PAYLOAD)
sleep(0.5) sleep(0.5)
self.assertNotEqual(mainSocket.namespace.payload, PAYLOAD) self.assertNotEqual(mainSocket._namespace.payload, PAYLOAD)
self.assertNotEqual(chatSocket.namespace.payload, PAYLOAD) self.assertNotEqual(chatSocket._namespace.payload, PAYLOAD)
self.assertEqual(newsSocket.namespace.payload, PAYLOAD) self.assertEqual(newsSocket._namespace.payload, PAYLOAD)
def test_delete(self):
socketIO = SocketIO('localhost', 8000)
childThreads = [
socketIO._heartbeatThread,
socketIO._namespaceThread,
]
del socketIO
for childThread in childThreads:
self.assertEqual(True, childThread.done.is_set())
class Namespace(BaseNamespace): class Namespace(BaseNamespace):
payload = None payload = None
def on_ddd(self, data): def on_ddd(self, data=''):
self.payload = data self.payload = data