diff --git a/mastodon/streaming.py b/mastodon/streaming.py index 0bb79fe..3f622b7 100644 --- a/mastodon/streaming.py +++ b/mastodon/streaming.py @@ -157,12 +157,12 @@ def handle_stream(self, response): raise exception from err except ReadTimeout as err: exception = MastodonReadTimeout( - "Timed out while reading from server."), + "Timed out while reading from server.") self.on_abort(exception) raise exception from err except ConnectionError as err: exception = MastodonNetworkError( - "Requests reports connection error."), + "Requests reports connection error.") self.on_abort(exception) raise exception from err diff --git a/tests/test_streaming.py b/tests/test_streaming.py index cc1eb8f..83eb149 100644 --- a/tests/test_streaming.py +++ b/tests/test_streaming.py @@ -1,8 +1,9 @@ import pytest import itertools from mastodon.streaming import StreamListener, CallbackStreamListener -from mastodon.Mastodon import MastodonMalformedEventError +from mastodon.Mastodon import MastodonMalformedEventError, MastodonNetworkError, MastodonReadTimeout from mastodon import Mastodon +from requests.exceptions import ChunkedEncodingError, ConnectionError, ReadTimeout import threading import time @@ -90,6 +91,10 @@ def __init__(self): self.heartbeats = 0 self.bla_called = False self.do_something_called = False + self.aborts = [] + + def on_abort(self, err): + self.aborts.append(err) def on_update(self, status): self.updates.append(status) @@ -345,6 +350,36 @@ def test_multiline_payload(): ]) assert listener.updates == [{"foo": "bar"}] + +class RaisingResponse(): + """A response object whose iter_content immediately raises `exception`.""" + def __init__(self, exception): + self.exception = exception + + def iter_content(self, chunk_size): + raise self.exception + yield # pragma: no cover - makes this a generator function + + +@pytest.mark.parametrize("raised,expected", [ + (ChunkedEncodingError("nope"), MastodonNetworkError), + (ReadTimeout("nope"), MastodonReadTimeout), + (ConnectionError("nope"), MastodonNetworkError), +]) +def test_handle_stream_transport_errors(raised, expected): + """ + Transport level errors have to surface as the matching Mastodon.py error, and the + same object has to be handed to on_abort. Everything that reaches the caller of + handle_stream must also be a MastodonNetworkError, because that is what the + auto-reconnect handler in Mastodon.__stream catches to decide to reconnect. + """ + listener = Listener() + with pytest.raises(expected) as exc_info: + listener.handle_stream(RaisingResponse(raised)) + assert isinstance(exc_info.value, MastodonNetworkError) + assert exc_info.value.__cause__ is raised + assert listener.aborts == [exc_info.value] + @pytest.mark.vcr(match_on=['path']) def test_stream_user_direct(api, api2, api3, vcr): patch_streaming()