Skip to content

Commit 0e0b26f

Browse files
Improved handling of rejected connections
1 parent 5b79e28 commit 0e0b26f

9 files changed

Lines changed: 75 additions & 15 deletions

File tree

docs/server.rst

Lines changed: 8 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -181,6 +181,14 @@ The ``sid`` argument passed into all the event handlers is a connection
181181
identifier for the client. All the events from a client will use the same
182182
``sid`` value.
183183

184+
The ``connect`` handler is the place where the server can perform
185+
authentication. The value returned by this handler is used to determine if the
186+
connection is accepted or rejected. When the handler does not return any value
187+
(which is the same as returning ``None``) or when it returns ``True`` the
188+
connection is accepted. If the handler returns ``False`` or any JSON
189+
compatible data type (string, integer, list or dictionary) the connection is
190+
rejected. A rejected connection triggers a response with a 401 status code.
191+
184192
The ``data`` argument passed to the ``'message'`` event handler contains
185193
application-specific data provided by the client with the event.
186194

engineio/asyncio_client.py

Lines changed: 2 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -183,9 +183,10 @@ async def _connect_polling(self, url, headers, engineio_path):
183183
raise exceptions.ConnectionError(
184184
'Connection refused by the server')
185185
if r.status < 200 or r.status >= 300:
186+
self._reset()
186187
raise exceptions.ConnectionError(
187188
'Unexpected status code {} in server response'.format(
188-
r.status))
189+
r.status), await r.json())
189190
try:
190191
p = payload.Payload(encoded_payload=await r.read())
191192
except ValueError:

engineio/asyncio_server.py

Lines changed: 2 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -384,10 +384,10 @@ async def _handle_connect(self, environ, transport, b64=False,
384384

385385
ret = await self._trigger_event('connect', sid, environ,
386386
run_async=False)
387-
if ret is False:
387+
if ret is not None and ret is not True:
388388
del self.sockets[sid]
389389
self.logger.warning('Application rejected connection')
390-
return self._unauthorized()
390+
return self._unauthorized(ret or None)
391391

392392
if transport == 'websocket':
393393
ret = await s.handle_get_request(environ)

engineio/client.py

Lines changed: 2 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -295,9 +295,10 @@ def _connect_polling(self, url, headers, engineio_path):
295295
raise exceptions.ConnectionError(
296296
'Connection refused by the server')
297297
if r.status_code < 200 or r.status_code >= 300:
298+
self._reset()
298299
raise exceptions.ConnectionError(
299300
'Unexpected status code {} in server response'.format(
300-
r.status_code))
301+
r.status_code), r.json())
301302
try:
302303
p = payload.Payload(encoded_payload=r.content)
303304
except ValueError:

engineio/server.py

Lines changed: 8 additions & 5 deletions
Original file line numberDiff line numberDiff line change
@@ -506,10 +506,10 @@ def _handle_connect(self, environ, start_response, transport, b64=False,
506506
s.send(pkt)
507507

508508
ret = self._trigger_event('connect', sid, environ, run_async=False)
509-
if ret is False:
509+
if ret is not None and ret is not True:
510510
del self.sockets[sid]
511511
self.logger.warning('Application rejected connection')
512-
return self._unauthorized()
512+
return self._unauthorized(ret or None)
513513

514514
if transport == 'websocket':
515515
ret = s.handle_get_request(environ, start_response)
@@ -592,11 +592,14 @@ def _method_not_found(self):
592592
'headers': [('Content-Type', 'text/plain')],
593593
'response': b'Method Not Found'}
594594

595-
def _unauthorized(self):
595+
def _unauthorized(self, message=None):
596596
"""Generate a unauthorized HTTP error response."""
597+
if message is None:
598+
message = 'Unauthorized'
599+
message = packet.Packet.json.dumps(message)
597600
return {'status': '401 UNAUTHORIZED',
598-
'headers': [('Content-Type', 'text/plain')],
599-
'response': b'Unauthorized'}
601+
'headers': [('Content-Type', 'application/json')],
602+
'response': message.encode('utf-8')}
600603

601604
def _cors_allowed_origins(self, environ):
602605
default_origins = []

tests/asyncio/test_asyncio_client.py

Lines changed: 9 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -272,8 +272,15 @@ def test_polling_connection_404(self):
272272
c = asyncio_client.AsyncClient()
273273
c._send_request = AsyncMock()
274274
c._send_request.mock.return_value.status = 404
275-
self.assertRaises(
276-
exceptions.ConnectionError, _run, c.connect('http://foo'))
275+
c._send_request.mock.return_value.json = AsyncMock(
276+
return_value={'foo': 'bar'})
277+
try:
278+
_run(c.connect('http://foo'))
279+
except exceptions.ConnectionError as exc:
280+
self.assertEqual(len(exc.args), 2)
281+
self.assertEqual(exc.args[0],
282+
'Unexpected status code 404 in server response')
283+
self.assertEqual(exc.args[1], {'foo': 'bar'})
277284

278285
def test_polling_connection_invalid_packet(self):
279286
c = asyncio_client.AsyncClient()

tests/asyncio/test_asyncio_server.py

Lines changed: 20 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -599,6 +599,26 @@ def mock_connect(sid, environ):
599599
self.assertEqual(len(s.sockets), 0)
600600
self.assertEqual(a._async['make_response'].call_args[0][0],
601601
'401 UNAUTHORIZED')
602+
self.assertEqual(a._async['make_response'].call_args[0][2],
603+
b'"Unauthorized"')
604+
605+
@mock.patch('importlib.import_module')
606+
def test_connect_event_rejects_with_message(self, import_module):
607+
a = self.get_async_mock()
608+
import_module.side_effect = [a]
609+
s = asyncio_server.AsyncServer()
610+
s._generate_id = mock.MagicMock(return_value='123')
611+
612+
def mock_connect(sid, environ):
613+
return {'not': 'allowed'}
614+
615+
s.on('connect')(mock_connect)
616+
_run(s.handle_request('request'))
617+
self.assertEqual(len(s.sockets), 0)
618+
self.assertEqual(a._async['make_response'].call_args[0][0],
619+
'401 UNAUTHORIZED')
620+
self.assertEqual(a._async['make_response'].call_args[0][2],
621+
b'{"not": "allowed"}')
602622

603623
@mock.patch('importlib.import_module')
604624
def test_method_not_found(self, import_module):

tests/common/test_client.py

Lines changed: 9 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -294,8 +294,15 @@ def test_polling_connection_failed(self, _send_request, _time):
294294
@mock.patch('engineio.client.Client._send_request')
295295
def test_polling_connection_404(self, _send_request):
296296
_send_request.return_value.status_code = 404
297-
c = client.Client()
298-
self.assertRaises(exceptions.ConnectionError, c.connect, 'http://foo')
297+
_send_request.return_value.json.return_value = {'foo': 'bar'}
298+
c = client.Client()
299+
try:
300+
c.connect('http://foo')
301+
except exceptions.ConnectionError as exc:
302+
self.assertEqual(len(exc.args), 2)
303+
self.assertEqual(exc.args[0],
304+
'Unexpected status code 404 in server response')
305+
self.assertEqual(exc.args[1], {'foo': 'bar'})
299306

300307
@mock.patch('engineio.client.Client._send_request')
301308
def test_polling_connection_invalid_packet(self, _send_request):

tests/common/test_server.py

Lines changed: 15 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -724,7 +724,7 @@ def test_connect_cors_disabled_no_origin(self):
724724
def test_connect_event(self):
725725
s = server.Server()
726726
s._generate_id = mock.MagicMock(return_value='123')
727-
mock_event = mock.MagicMock()
727+
mock_event = mock.MagicMock(return_value=None)
728728
s.on('connect')(mock_event)
729729
environ = {'REQUEST_METHOD': 'GET', 'QUERY_STRING': ''}
730730
start_response = mock.MagicMock()
@@ -739,9 +739,22 @@ def test_connect_event_rejects(self):
739739
s.on('connect')(mock_event)
740740
environ = {'REQUEST_METHOD': 'GET', 'QUERY_STRING': ''}
741741
start_response = mock.MagicMock()
742-
s.handle_request(environ, start_response)
742+
ret = s.handle_request(environ, start_response)
743+
self.assertEqual(len(s.sockets), 0)
744+
self.assertEqual(start_response.call_args[0][0], '401 UNAUTHORIZED')
745+
self.assertEqual(ret, [b'"Unauthorized"'])
746+
747+
def test_connect_event_rejects_with_message(self):
748+
s = server.Server()
749+
s._generate_id = mock.MagicMock(return_value='123')
750+
mock_event = mock.MagicMock(return_value='not allowed')
751+
s.on('connect')(mock_event)
752+
environ = {'REQUEST_METHOD': 'GET', 'QUERY_STRING': ''}
753+
start_response = mock.MagicMock()
754+
ret = s.handle_request(environ, start_response)
743755
self.assertEqual(len(s.sockets), 0)
744756
self.assertEqual(start_response.call_args[0][0], '401 UNAUTHORIZED')
757+
self.assertEqual(ret, [b'"not allowed"'])
745758

746759
def test_method_not_found(self):
747760
s = server.Server()

0 commit comments

Comments
 (0)