diff --git a/docs/source/CHANGES.rst b/docs/source/CHANGES.rst index 944efaa..0edbefd 100644 --- a/docs/source/CHANGES.rst +++ b/docs/source/CHANGES.rst @@ -2,6 +2,12 @@ pystalk ChangeLog ################# +===== +0.9.1 +===== +* Raise :py:class:`pystalk.client.BeanstalkConnectionError` when an established + connection is closed by beanstalkd. + ===== 0.9.0 ===== diff --git a/pystalk/__init__.py b/pystalk/__init__.py index 330754e..61a885b 100644 --- a/pystalk/__init__.py +++ b/pystalk/__init__.py @@ -1,7 +1,7 @@ from .client import BeanstalkClient, BeanstalkError, BeanstalkConnectionError from .pool import ProductionPool -__version__ = '0.9.0' +__version__ = '0.9.1' __author__ = 'EasyPost ' diff --git a/pystalk/client.py b/pystalk/client.py index 2f99397..58820bf 100644 --- a/pystalk/client.py +++ b/pystalk/client.py @@ -225,13 +225,18 @@ def close(self): def _sock_ctx(self): yield self._socket + def _receive_from_socket(self, sock, size): + message = sock.recv(size) + if not message: + connection_error = ConnectionResetError('Connection closed by beanstalkd') + raise BeanstalkConnectionError(self.host, self.port, connection_error) from connection_error + return message + def _receive_data_with_prefix(self, prefix, sock): buf = b'' target_len = len(prefix) + 28 - while b'\r\n' not in buf: - message = sock.recv(target_len - len(buf)) - if not message: - break + while b'\r\n' not in buf and len(buf) < target_len: + message = self._receive_from_socket(sock, target_len - len(buf)) buf += message if b' ' not in buf: error = buf.rstrip() @@ -244,10 +249,8 @@ def _receive_data_with_prefix(self, prefix, sock): def _receive_id_and_data_with_prefix(self, prefix, sock): buf = b'' target_len = len(prefix) + 28 - while b'\r\n' not in buf: - message = sock.recv(target_len - len(buf)) - if not message: - break + while b'\r\n' not in buf and len(buf) < target_len: + message = self._receive_from_socket(sock, target_len - len(buf)) buf += message if b' ' not in buf: error = buf.rstrip() @@ -260,15 +263,13 @@ def _receive_id_and_data_with_prefix(self, prefix, sock): def _receive_data(self, sock, initial=None): if initial is None: - initial = sock.recv(12) + initial = self._receive_from_socket(sock, 12) byte_length, rest = initial.split(b'\r\n', 1) byte_length = int(byte_length) + 2 buf = [rest] bytes_read = len(rest) while bytes_read < byte_length: - message = sock.recv(min(4096, byte_length - bytes_read)) - if not message: - break + message = self._receive_from_socket(sock, min(4096, byte_length - bytes_read)) bytes_read += len(message) buf.append(message) bytez = b''.join(buf)[:-2] @@ -282,7 +283,7 @@ def _receive_id(self, sock): return status, int(gid) def _receive_name(self, sock): - message = sock.recv(1024) + message = self._receive_from_socket(sock, 1024) if b' ' in message: status, rest = message.split(b' ', 1) return status, rest.rstrip() @@ -290,7 +291,7 @@ def _receive_name(self, sock): raise BeanstalkError(message.rstrip()) def _receive_word(self, sock, *expected_words): - message = sock.recv(1024).rstrip() + message = self._receive_from_socket(sock, 1024).rstrip() if message not in expected_words: raise BeanstalkError(message) return message diff --git a/tests/unit/test_pystalk.py b/tests/unit/test_pystalk.py index 097eb87..ca72a67 100644 --- a/tests/unit/test_pystalk.py +++ b/tests/unit/test_pystalk.py @@ -41,6 +41,17 @@ def test_stats(client, server): assert server.received == [b'stats\r\n'] +def test_reserve_raises_connection_error_when_server_closes_connection(client, server): + server.responses.append(b'') + + with pytest.raises(pystalk.BeanstalkConnectionError) as exc_info: + client.reserve_job(0) + + assert isinstance(exc_info.value.err, ConnectionResetError) + assert exc_info.value.host == 'pystalk.example.com' + assert exc_info.value.port == 0 + + def test_put_job_uses_utf8_byte_length_for_non_ascii(client, server): server.responses.append(b'INSERTED 1\r\n')