Added support for acks thanks to zratic
This commit is contained in:
parent
61918597b7
commit
18d2a1ea8c
2 changed files with 138 additions and 109 deletions
|
|
@ -12,9 +12,9 @@ PROTOCOL = 1 # socket.io protocol version
|
||||||
class BaseNamespace(object): # pragma: no cover
|
class BaseNamespace(object): # pragma: no cover
|
||||||
'Define socket.io behavior'
|
'Define socket.io behavior'
|
||||||
|
|
||||||
def __init__(self, _socketIO, namespacePath):
|
def __init__(self, _socketIO, path):
|
||||||
self._socketIO = _socketIO
|
self._socketIO = _socketIO
|
||||||
self._namespacePath = namespacePath
|
self._path = path
|
||||||
self._callbackByEvent = {}
|
self._callbackByEvent = {}
|
||||||
|
|
||||||
def on_connect(self):
|
def on_connect(self):
|
||||||
|
|
@ -26,11 +26,16 @@ class BaseNamespace(object): # pragma: no cover
|
||||||
def on_error(self, reason, advice):
|
def on_error(self, reason, advice):
|
||||||
print '[Error] %s' % advice
|
print '[Error] %s' % advice
|
||||||
|
|
||||||
def on_message(self, messageData):
|
def on_message(self, data):
|
||||||
print '[Message] %s' % messageData
|
print '[Message] %s' % data
|
||||||
|
|
||||||
def on_default(self, eventName, *eventArguments):
|
def on_default(self, event, *args):
|
||||||
print '[Event] %s%s' % (eventName, eventArguments)
|
callback, args = find_callback(args)
|
||||||
|
arguments = [str(_) for _ in args]
|
||||||
|
if callback:
|
||||||
|
arguments.append('callback(*args)')
|
||||||
|
callback()
|
||||||
|
print '[Event] %s(%s)' % (event, ', '.join(arguments))
|
||||||
|
|
||||||
def on_open(self, *args):
|
def on_open(self, *args):
|
||||||
print '[Open]', args
|
print '[Open]', args
|
||||||
|
|
@ -44,28 +49,25 @@ class BaseNamespace(object): # pragma: no cover
|
||||||
def on_reconnect(self, *args):
|
def on_reconnect(self, *args):
|
||||||
print '[Reconnect]', args
|
print '[Reconnect]', args
|
||||||
|
|
||||||
def message(self, messageData, messageCallback=None):
|
def message(self, data, callback=None):
|
||||||
self._socketIO.message(
|
self._socketIO.message(data, callback, path=self._path)
|
||||||
messageData, messageCallback, namespacePath=self._namespacePath)
|
|
||||||
|
|
||||||
def emit(self, eventName, *eventArguments):
|
def emit(self, event, *args, **kw):
|
||||||
self._socketIO.emit(
|
kw['path'] = self._path
|
||||||
eventName, *eventArguments, namespacePath=self._namespacePath)
|
self._socketIO.emit(event, *args, **kw)
|
||||||
|
|
||||||
def on(self, eventName, eventCallback):
|
def on(self, event, callback):
|
||||||
self._callbackByEvent[eventName] = eventCallback
|
self._callbackByEvent[event] = callback
|
||||||
|
|
||||||
def _get_eventCallback(self, eventName):
|
def _get_eventCallback(self, event):
|
||||||
# Check callbacks defined by on()
|
# Check callbacks defined by on()
|
||||||
try:
|
try:
|
||||||
return self._callbackByEvent[eventName]
|
return self._callbackByEvent[event]
|
||||||
except KeyError:
|
except KeyError:
|
||||||
pass
|
pass
|
||||||
|
|
||||||
# Check callbacks defined explicitly or use on_default()
|
# Check callbacks defined explicitly or use on_default()
|
||||||
def callback(*eventArguments):
|
callback = lambda *args: self.on_default(event, *args)
|
||||||
return self.on_default(eventName, *eventArguments)
|
return getattr(self, 'on_' + event.replace(' ', '_'), callback)
|
||||||
return getattr(self, 'on_' + eventName.replace(' ', '_'), callback)
|
|
||||||
|
|
||||||
|
|
||||||
class SocketIO(object):
|
class SocketIO(object):
|
||||||
|
|
@ -99,38 +101,38 @@ class SocketIO(object):
|
||||||
self.disconnect()
|
self.disconnect()
|
||||||
|
|
||||||
def __del__(self):
|
def __del__(self):
|
||||||
self.disconnect(closeSocket=False)
|
self.disconnect(close=False)
|
||||||
|
|
||||||
@property
|
@property
|
||||||
def connected(self):
|
def connected(self):
|
||||||
return self._socketIO.connected
|
return self._socketIO.connected
|
||||||
|
|
||||||
def disconnect(self, namespacePath='', closeSocket=True):
|
def disconnect(self, path='', close=True):
|
||||||
if self.connected:
|
if self.connected:
|
||||||
self._socketIO.disconnect(namespacePath, closeSocket)
|
self._socketIO.disconnect(path, close)
|
||||||
if namespacePath:
|
if path:
|
||||||
del self._namespaceByPath[namespacePath]
|
del self._namespaceByPath[path]
|
||||||
else:
|
else:
|
||||||
self._rhythmicThread.cancel()
|
self._rhythmicThread.cancel()
|
||||||
self._listenerThread.cancel()
|
self._listenerThread.cancel()
|
||||||
|
|
||||||
def define(self, Namespace, namespacePath=''):
|
def define(self, Namespace, path=''):
|
||||||
self._socketIO.connect(namespacePath)
|
self._socketIO.connect(path)
|
||||||
namespace = Namespace(self._socketIO, namespacePath)
|
namespace = Namespace(self._socketIO, path)
|
||||||
self._namespaceByPath[namespacePath] = namespace
|
self._namespaceByPath[path] = namespace
|
||||||
return namespace
|
return namespace
|
||||||
|
|
||||||
def get_namespace(self, namespacePath=''):
|
def get_namespace(self, path=''):
|
||||||
return self._namespaceByPath[namespacePath]
|
return self._namespaceByPath[path]
|
||||||
|
|
||||||
def on(self, eventName, eventCallback, namespacePath=''):
|
def on(self, event, callback, path=''):
|
||||||
return self.get_namespace(namespacePath).on(eventName, eventCallback)
|
return self.get_namespace(path).on(event, callback)
|
||||||
|
|
||||||
def message(self, messageData, messageCallback=None, namespacePath=''):
|
def message(self, data, callback=None, path=''):
|
||||||
self._socketIO.message(messageData, messageCallback, namespacePath)
|
self._socketIO.message(data, callback, path)
|
||||||
|
|
||||||
def emit(self, eventName, *eventArguments, **eventKeywords):
|
def emit(self, event, *args, **kw):
|
||||||
self._socketIO.emit(eventName, *eventArguments, **eventKeywords)
|
self._socketIO.emit(event, *args, **kw)
|
||||||
|
|
||||||
def wait(self, seconds=None, forCallbacks=False):
|
def wait(self, seconds=None, forCallbacks=False):
|
||||||
if forCallbacks:
|
if forCallbacks:
|
||||||
|
|
@ -187,10 +189,13 @@ class _ListenerThread(Thread):
|
||||||
# Block callingThread until listenerThread terminates
|
# Block callingThread until listenerThread terminates
|
||||||
self.join(seconds)
|
self.join(seconds)
|
||||||
|
|
||||||
|
def get_ackCallback(self, packetID):
|
||||||
|
return lambda *args: self._socketIO.ack(packetID, *args)
|
||||||
|
|
||||||
def run(self):
|
def run(self):
|
||||||
while not self.done.is_set():
|
while not self.done.is_set():
|
||||||
try:
|
try:
|
||||||
code, packetID, namespacePath, data = self._socketIO.recv_packet()
|
code, packetID, path, data = self._socketIO.recv_packet()
|
||||||
except SocketIOConnectionError, error:
|
except SocketIOConnectionError, error:
|
||||||
print error
|
print error
|
||||||
return
|
return
|
||||||
|
|
@ -198,9 +203,9 @@ class _ListenerThread(Thread):
|
||||||
print error
|
print error
|
||||||
continue
|
continue
|
||||||
try:
|
try:
|
||||||
namespace = self._namespaceByPath[namespacePath]
|
namespace = self._namespaceByPath[path]
|
||||||
except KeyError:
|
except KeyError:
|
||||||
print 'Received unexpected namespacePath (%s)' % namespacePath
|
print 'Received unexpected path (%s)' % path
|
||||||
continue
|
continue
|
||||||
try:
|
try:
|
||||||
delegate = {
|
delegate = {
|
||||||
|
|
@ -228,25 +233,33 @@ class _ListenerThread(Thread):
|
||||||
pass
|
pass
|
||||||
|
|
||||||
def on_message(self, packetID, get_eventCallback, data):
|
def on_message(self, packetID, get_eventCallback, data):
|
||||||
get_eventCallback('message')(data)
|
args = [data]
|
||||||
|
if packetID:
|
||||||
|
args.append(self.get_ackCallback(packetID))
|
||||||
|
get_eventCallback('message')(args)
|
||||||
|
|
||||||
def on_json(self, packetID, get_eventCallback, data):
|
def on_json(self, packetID, get_eventCallback, data):
|
||||||
get_eventCallback('message')(loads(data))
|
args = [loads(data)]
|
||||||
|
if packetID:
|
||||||
|
args.append(self.get_ackCallback(packetID))
|
||||||
|
get_eventCallback('message')(args)
|
||||||
|
|
||||||
def on_event(self, packetID, get_eventCallback, data):
|
def on_event(self, packetID, get_eventCallback, data):
|
||||||
valueByName = loads(data)
|
valueByName = loads(data)
|
||||||
eventName = valueByName['name']
|
event = valueByName['name']
|
||||||
eventArguments = valueByName.get('args', [])
|
args = valueByName.get('args', [])
|
||||||
get_eventCallback(eventName)(*eventArguments)
|
if packetID:
|
||||||
|
args.append(self.get_ackCallback(packetID))
|
||||||
|
get_eventCallback(event)(*args)
|
||||||
|
|
||||||
def on_acknowledgment(self, packetID, get_eventCallback, data):
|
def on_acknowledgment(self, packetID, get_eventCallback, data):
|
||||||
dataParts = data.split('+', 1)
|
dataParts = data.split('+', 1)
|
||||||
messageID = int(dataParts[0])
|
messageID = int(dataParts[0])
|
||||||
arguments = loads(dataParts[1]) or []
|
args = loads(dataParts[1]) or []
|
||||||
messageCallback = self._socketIO.get_messageCallback(messageID)
|
callback = self._socketIO.get_messageCallback(messageID)
|
||||||
if not messageCallback:
|
if not callback:
|
||||||
return
|
return
|
||||||
messageCallback(*arguments)
|
callback(*args)
|
||||||
if self.waiting.is_set() and not self._socketIO.has_messageCallback:
|
if self.waiting.is_set() and not self._socketIO.has_messageCallback:
|
||||||
self.cancel()
|
self.cancel()
|
||||||
|
|
||||||
|
|
@ -284,18 +297,18 @@ class _SocketIO(object):
|
||||||
self.callbackByMessageID = {}
|
self.callbackByMessageID = {}
|
||||||
|
|
||||||
def __del__(self):
|
def __del__(self):
|
||||||
self.disconnect(closeSocket=False)
|
self.disconnect(close=False)
|
||||||
|
|
||||||
def disconnect(self, namespacePath='', closeSocket=True):
|
def disconnect(self, path='', close=True):
|
||||||
if not self.connected:
|
if not self.connected:
|
||||||
return
|
return
|
||||||
if namespacePath:
|
if path:
|
||||||
self.send_packet(0, namespacePath)
|
self.send_packet(0, path)
|
||||||
elif closeSocket:
|
elif close:
|
||||||
self.connection.close()
|
self.connection.close()
|
||||||
|
|
||||||
def connect(self, namespacePath):
|
def connect(self, path):
|
||||||
self.send_packet(1, namespacePath)
|
self.send_packet(1, path)
|
||||||
|
|
||||||
def send_heartbeat(self):
|
def send_heartbeat(self):
|
||||||
try:
|
try:
|
||||||
|
|
@ -304,24 +317,25 @@ class _SocketIO(object):
|
||||||
print 'Could not send heartbeat'
|
print 'Could not send heartbeat'
|
||||||
pass
|
pass
|
||||||
|
|
||||||
def message(self, messageData, messageCallback, namespacePath):
|
def message(self, data, callback, path):
|
||||||
if isinstance(messageData, basestring):
|
if isinstance(data, basestring):
|
||||||
code = 3
|
code = 3
|
||||||
data = messageData
|
packetData = data
|
||||||
else:
|
else:
|
||||||
code = 4
|
code = 4
|
||||||
data = dumps(messageData, ensure_ascii=False)
|
packetData = dumps(data, ensure_ascii=False)
|
||||||
self.send_packet(code, namespacePath, data, messageCallback)
|
self.send_packet(code, path, packetData, callback)
|
||||||
|
|
||||||
def emit(self, eventName, *eventArguments, **eventKeywords):
|
def emit(self, event, *args, **kw):
|
||||||
if eventArguments and callable(eventArguments[-1]):
|
callback, args = find_callback(args, kw)
|
||||||
messageCallback = eventArguments[-1]
|
packetData = dumps(dict(name=event, args=args), ensure_ascii=False)
|
||||||
eventArguments = eventArguments[:-1]
|
path = kw.get('path', '')
|
||||||
else:
|
self.send_packet(5, path, packetData, callback)
|
||||||
messageCallback = None
|
|
||||||
namespacePath = eventKeywords.get('namespacePath', '')
|
def ack(self, packetID, *args):
|
||||||
data = dumps(dict(name=eventName, args=eventArguments), ensure_ascii=False)
|
packetID = packetID.rstrip('+')
|
||||||
self.send_packet(5, namespacePath, data, messageCallback)
|
packetData = '%s+%s' % (packetID, dumps(args, ensure_ascii=False)) if args else packetID
|
||||||
|
self.send_packet(6, data=packetData)
|
||||||
|
|
||||||
def set_messageCallback(self, callback):
|
def set_messageCallback(self, callback):
|
||||||
'Set callback that will be called after receiving an acknowledgment'
|
'Set callback that will be called after receiving an acknowledgment'
|
||||||
|
|
@ -355,18 +369,18 @@ class _SocketIO(object):
|
||||||
except AttributeError:
|
except AttributeError:
|
||||||
raise SocketIOPacketError('Received invalid packet (%s)' % packet)
|
raise SocketIOPacketError('Received invalid packet (%s)' % packet)
|
||||||
packetCount = len(packetParts)
|
packetCount = len(packetParts)
|
||||||
code, packetID, namespacePath, data = None, None, None, None
|
code, packetID, path, data = None, None, None, None
|
||||||
if 4 == packetCount:
|
if 4 == packetCount:
|
||||||
code, packetID, namespacePath, data = packetParts
|
code, packetID, path, data = packetParts
|
||||||
elif 3 == packetCount:
|
elif 3 == packetCount:
|
||||||
code, packetID, namespacePath = packetParts
|
code, packetID, path = packetParts
|
||||||
elif 1 == packetCount:
|
elif 1 == packetCount:
|
||||||
code = packetParts[0]
|
code = packetParts[0]
|
||||||
return code, packetID, namespacePath, data
|
return code, packetID, path, data
|
||||||
|
|
||||||
def send_packet(self, code, namespacePath='', data='', messageCallback=None):
|
def send_packet(self, code, path='', data='', callback=None):
|
||||||
callbackNumber = self.set_messageCallback(messageCallback) if messageCallback else ''
|
packetID = self.set_messageCallback(callback) if callback else ''
|
||||||
packetParts = [str(code), callbackNumber, namespacePath, data]
|
packetParts = [str(code), packetID, path, data]
|
||||||
try:
|
try:
|
||||||
self.connection.send(':'.join(packetParts))
|
self.connection.send(':'.join(packetParts))
|
||||||
except socket.error:
|
except socket.error:
|
||||||
|
|
@ -387,3 +401,13 @@ class SocketIOConnectionError(SocketIOError):
|
||||||
|
|
||||||
class SocketIOPacketError(SocketIOError):
|
class SocketIOPacketError(SocketIOError):
|
||||||
pass
|
pass
|
||||||
|
|
||||||
|
|
||||||
|
def find_callback(args, kw=None):
|
||||||
|
'Return callback whether passed as a last argument or as a keyword'
|
||||||
|
if args and callable(args[-1]):
|
||||||
|
return args[-1], args[:-1]
|
||||||
|
try:
|
||||||
|
return kw['callback'], args
|
||||||
|
except (KeyError, TypeError):
|
||||||
|
return None, args
|
||||||
|
|
|
||||||
|
|
@ -1,9 +1,8 @@
|
||||||
from socketIO_client import SocketIO, BaseNamespace
|
from socketIO_client import SocketIO, BaseNamespace, find_callback
|
||||||
from time import sleep
|
from time import sleep
|
||||||
from unittest import TestCase
|
from unittest import TestCase
|
||||||
|
|
||||||
|
|
||||||
ON_RESPONSE_CALLED = False
|
|
||||||
PORT = 8000
|
PORT = 8000
|
||||||
PAYLOAD = {'xxx': 'yyy'}
|
PAYLOAD = {'xxx': 'yyy'}
|
||||||
|
|
||||||
|
|
@ -11,13 +10,28 @@ PAYLOAD = {'xxx': 'yyy'}
|
||||||
class TestSocketIO(TestCase):
|
class TestSocketIO(TestCase):
|
||||||
|
|
||||||
def setUp(self):
|
def setUp(self):
|
||||||
global ON_RESPONSE_CALLED
|
|
||||||
ON_RESPONSE_CALLED = False
|
|
||||||
self.socketIO = SocketIO('localhost', PORT)
|
self.socketIO = SocketIO('localhost', PORT)
|
||||||
|
self.called_on_response = False
|
||||||
|
|
||||||
def tearDown(self):
|
def tearDown(self):
|
||||||
del self.socketIO
|
del self.socketIO
|
||||||
|
|
||||||
|
def on_response(self, *args):
|
||||||
|
self.called_on_response = True
|
||||||
|
callback, args = find_callback(args)
|
||||||
|
if callback:
|
||||||
|
callback(*args)
|
||||||
|
|
||||||
|
def test_disconnect(self):
|
||||||
|
childThreads = [
|
||||||
|
self.socketIO._rhythmicThread,
|
||||||
|
self.socketIO._listenerThread,
|
||||||
|
]
|
||||||
|
self.socketIO.disconnect()
|
||||||
|
for childThread in childThreads:
|
||||||
|
self.assertEqual(True, childThread.done.is_set())
|
||||||
|
self.assertEqual(False, self.socketIO.connected)
|
||||||
|
|
||||||
def test_emit(self):
|
def test_emit(self):
|
||||||
self.socketIO.define(Namespace)
|
self.socketIO.define(Namespace)
|
||||||
self.socketIO.emit('aaa')
|
self.socketIO.emit('aaa')
|
||||||
|
|
@ -31,20 +45,26 @@ class TestSocketIO(TestCase):
|
||||||
self.assertEqual(self.socketIO.get_namespace().payload, PAYLOAD)
|
self.assertEqual(self.socketIO.get_namespace().payload, PAYLOAD)
|
||||||
|
|
||||||
def test_emit_with_callback(self):
|
def test_emit_with_callback(self):
|
||||||
self.socketIO.emit('aaa', PAYLOAD, on_response)
|
self.socketIO.emit('aaa', PAYLOAD, self.on_response)
|
||||||
self.socketIO.wait(forCallbacks=True)
|
self.socketIO.wait(seconds=0.1, forCallbacks=True)
|
||||||
self.assertEqual(ON_RESPONSE_CALLED, True)
|
self.assertEqual(self.called_on_response, True)
|
||||||
|
|
||||||
def test_message(self):
|
def test_emit_with_event(self):
|
||||||
self.socketIO.message(PAYLOAD, on_response)
|
self.socketIO.on('aaa_response', self.on_response)
|
||||||
self.socketIO.wait(forCallbacks=True)
|
|
||||||
self.assertEqual(ON_RESPONSE_CALLED, True)
|
|
||||||
|
|
||||||
def test_events(self):
|
|
||||||
self.socketIO.on('aaa_response', on_response)
|
|
||||||
self.socketIO.emit('aaa', PAYLOAD)
|
self.socketIO.emit('aaa', PAYLOAD)
|
||||||
sleep(0.1)
|
sleep(0.1)
|
||||||
self.assertEqual(ON_RESPONSE_CALLED, True)
|
self.assertEqual(self.called_on_response, True)
|
||||||
|
|
||||||
|
def test_message(self):
|
||||||
|
self.socketIO.message(PAYLOAD, self.on_response)
|
||||||
|
self.socketIO.wait(seconds=0.1, forCallbacks=True)
|
||||||
|
self.assertEqual(self.called_on_response, True)
|
||||||
|
|
||||||
|
def test_ack(self):
|
||||||
|
self.socketIO.on('bbb_response', self.on_response)
|
||||||
|
self.socketIO.emit('bbb', PAYLOAD)
|
||||||
|
sleep(0.1)
|
||||||
|
self.assertEqual(self.called_on_response, True)
|
||||||
|
|
||||||
def test_namespaces(self):
|
def test_namespaces(self):
|
||||||
mainNamespace = self.socketIO.define(Namespace)
|
mainNamespace = self.socketIO.define(Namespace)
|
||||||
|
|
@ -59,19 +79,9 @@ class TestSocketIO(TestCase):
|
||||||
|
|
||||||
def test_namespaces_with_callback(self):
|
def test_namespaces_with_callback(self):
|
||||||
mainNamespace = self.socketIO.get_namespace()
|
mainNamespace = self.socketIO.get_namespace()
|
||||||
mainNamespace.message(PAYLOAD, on_response)
|
mainNamespace.message(PAYLOAD, self.on_response)
|
||||||
sleep(0.1)
|
sleep(0.1)
|
||||||
self.assertEqual(ON_RESPONSE_CALLED, True)
|
self.assertEqual(self.called_on_response, True)
|
||||||
|
|
||||||
def test_disconnect(self):
|
|
||||||
childThreads = [
|
|
||||||
self.socketIO._rhythmicThread,
|
|
||||||
self.socketIO._listenerThread,
|
|
||||||
]
|
|
||||||
self.socketIO.disconnect()
|
|
||||||
for childThread in childThreads:
|
|
||||||
self.assertEqual(True, childThread.done.is_set())
|
|
||||||
self.assertEqual(False, self.socketIO.connected)
|
|
||||||
|
|
||||||
|
|
||||||
class Namespace(BaseNamespace):
|
class Namespace(BaseNamespace):
|
||||||
|
|
@ -81,8 +91,3 @@ class Namespace(BaseNamespace):
|
||||||
def on_aaa_response(self, data=''):
|
def on_aaa_response(self, data=''):
|
||||||
print '[Event] aaa_response(%s)' % data
|
print '[Event] aaa_response(%s)' % data
|
||||||
self.payload = data
|
self.payload = data
|
||||||
|
|
||||||
|
|
||||||
def on_response(*args):
|
|
||||||
global ON_RESPONSE_CALLED
|
|
||||||
ON_RESPONSE_CALLED = True
|
|
||||||
|
|
|
||||||
Loading…
Add table
Add a link
Reference in a new issue