diff options
| author | Nikolay Kim <fafhrd91@gmail.com> | 2013-03-20 11:25:20 -0700 |
|---|---|---|
| committer | Nikolay Kim <fafhrd91@gmail.com> | 2013-03-20 11:25:20 -0700 |
| commit | 35981988fbf103d88d6415be790dfd2aa140260d (patch) | |
| tree | 322ce1a7bbb54544576429688b7dd713d230dd1b | |
| parent | b780325d12b2b0faca58d87792aa26f1067963cc (diff) | |
| download | trollius-git-35981988fbf103d88d6415be790dfd2aa140260d.tar.gz | |
timeout support for Future
| -rw-r--r-- | tests/base_events_test.py | 4 | ||||
| -rw-r--r-- | tests/events_test.py | 48 | ||||
| -rw-r--r-- | tests/tasks_test.py | 33 | ||||
| -rw-r--r-- | tulip/base_events.py | 4 | ||||
| -rw-r--r-- | tulip/futures.py | 13 | ||||
| -rw-r--r-- | tulip/locks.py | 36 | ||||
| -rw-r--r-- | tulip/tasks.py | 16 |
7 files changed, 86 insertions, 68 deletions
diff --git a/tests/base_events_test.py b/tests/base_events_test.py index 03c3296..1839837 100644 --- a/tests/base_events_test.py +++ b/tests/base_events_test.py @@ -11,6 +11,7 @@ from tulip import base_events from tulip import events from tulip import futures from tulip import protocols +from tulip import tasks from tulip import test_utils @@ -265,7 +266,8 @@ class BaseEventLoopTests(test_utils.LogTrackingTestCase): self.event_loop.getaddrinfo = getaddrinfo - task = self.event_loop.create_connection(MyProto, 'xkcd.com', 80) + task = tasks.Task( + self.event_loop.create_connection(MyProto, 'xkcd.com', 80)) task._step() exc = task.exception() self.assertEqual("Multiple exceptions: err1, err2", str(exc)) diff --git a/tests/events_test.py b/tests/events_test.py index 20c398d..1c0883b 100644 --- a/tests/events_test.py +++ b/tests/events_test.py @@ -25,6 +25,7 @@ from tulip import events from tulip import transports from tulip import protocols from tulip import selector_events +from tulip import tasks from tulip import test_utils @@ -531,7 +532,8 @@ class EventLoopTestsMixin: def test_create_connection(self): with self.run_test_server() as httpd: host, port = httpd.socket.getsockname() - f = self.event_loop.create_connection(MyProto, host, port) + f = tasks.Task( + self.event_loop.create_connection(MyProto, host, port)) tr, pr = self.event_loop.run_until_complete(f) self.assertTrue(isinstance(tr, transports.Transport)) self.assertTrue(isinstance(pr, protocols.Protocol)) @@ -541,7 +543,8 @@ class EventLoopTestsMixin: def test_create_connection_sock(self): with self.run_test_server() as httpd: host, port = httpd.socket.getsockname() - f = self.event_loop.create_connection(MyProto, host, port) + f = tasks.Task( + self.event_loop.create_connection(MyProto, host, port)) tr, pr = self.event_loop.run_until_complete(f) self.assertTrue(isinstance(tr, transports.Transport)) self.assertTrue(isinstance(pr, protocols.Protocol)) @@ -552,25 +555,26 @@ class EventLoopTestsMixin: def test_create_ssl_connection(self): with self.run_test_server(use_ssl=True) as httpsd: host, port = httpsd.socket.getsockname() - f = self.event_loop.create_connection( - MyProto, host, port, ssl=True) + f = tasks.Task(self.event_loop.create_connection( + MyProto, host, port, ssl=True)) tr, pr = self.event_loop.run_until_complete(f) self.assertTrue(isinstance(tr, transports.Transport)) self.assertTrue(isinstance(pr, protocols.Protocol)) self.assertTrue('ssl' in tr.__class__.__name__.lower()) - self.assertTrue(hasattr(tr.get_extra_info('socket'), 'getsockname')) + self.assertTrue( + hasattr(tr.get_extra_info('socket'), 'getsockname')) self.event_loop.run() self.assertTrue(pr.nbytes > 0) def test_create_connection_host_port_sock(self): self.suppress_log_errors() - fut = self.event_loop.create_connection( - MyProto, 'xkcd.com', 80, sock=object()) + fut = tasks.Task(self.event_loop.create_connection( + MyProto, 'xkcd.com', 80, sock=object())) self.assertRaises(ValueError, self.event_loop.run_until_complete, fut) def test_create_connection_no_host_port_sock(self): self.suppress_log_errors() - fut = self.event_loop.create_connection(MyProto) + fut = tasks.Task(self.event_loop.create_connection(MyProto)) self.assertRaises(ValueError, self.event_loop.run_until_complete, fut) def test_create_connection_no_getaddrinfo(self): @@ -578,7 +582,8 @@ class EventLoopTestsMixin: getaddrinfo = self.event_loop.getaddrinfo = unittest.mock.Mock() getaddrinfo.return_value = [] - fut = self.event_loop.create_connection(MyProto, 'xkcd.com', 80) + fut = tasks.Task( + self.event_loop.create_connection(MyProto, 'xkcd.com', 80)) self.assertRaises( socket.error, self.event_loop.run_until_complete, fut) @@ -587,7 +592,8 @@ class EventLoopTestsMixin: self.event_loop.sock_connect = unittest.mock.Mock() self.event_loop.sock_connect.side_effect = socket.error - fut = self.event_loop.create_connection(MyProto, 'xkcd.com', 80) + fut = tasks.Task( + self.event_loop.create_connection(MyProto, 'xkcd.com', 80)) self.assertRaises( socket.error, self.event_loop.run_until_complete, fut) @@ -602,7 +608,8 @@ class EventLoopTestsMixin: self.event_loop.sock_connect = unittest.mock.Mock() self.event_loop.sock_connect.side_effect = socket.error - fut = self.event_loop.create_connection(MyProto, 'xkcd.com', 80) + fut = tasks.Task( + self.event_loop.create_connection(MyProto, 'xkcd.com', 80)) self.assertRaises( socket.error, self.event_loop.run_until_complete, fut) @@ -721,8 +728,8 @@ class EventLoopTestsMixin: sock = self.event_loop.run_until_complete(f) host, port = sock.getsockname() - f = self.event_loop.create_datagram_connection( - MyDatagramProto, host, port) + f = tasks.Task(self.event_loop.create_datagram_connection( + MyDatagramProto, host, port)) transport, protocol = self.event_loop.run_until_complete(f) self.assertEqual('INITIALIZED', protocol.state) @@ -767,7 +774,8 @@ class EventLoopTestsMixin: sock = self.event_loop.run_until_complete(f) host, port = sock.getsockname() - f = self.event_loop.create_datagram_connection(MyDatagramProto) + f = tasks.Task( + self.event_loop.create_datagram_connection(MyDatagramProto)) transport, protocol = self.event_loop.run_until_complete(f) self.assertEqual('INITIALIZED', protocol.state) @@ -800,8 +808,8 @@ class EventLoopTestsMixin: getaddrinfo = self.event_loop.getaddrinfo = unittest.mock.Mock() getaddrinfo.return_value = [] - fut = self.event_loop.create_datagram_connection( - protocols.DatagramProtocol, 'xkcd.com', 80) + fut = tasks.Task(self.event_loop.create_datagram_connection( + protocols.DatagramProtocol, 'xkcd.com', 80)) self.assertRaises( socket.error, self.event_loop.run_until_complete, fut) @@ -810,8 +818,8 @@ class EventLoopTestsMixin: self.event_loop.sock_connect = unittest.mock.Mock() self.event_loop.sock_connect.side_effect = socket.error - fut = self.event_loop.create_datagram_connection( - protocols.DatagramProtocol, 'xkcd.com', 80) + fut = tasks.Task(self.event_loop.create_datagram_connection( + protocols.DatagramProtocol, 'xkcd.com', 80)) self.assertRaises( socket.error, self.event_loop.run_until_complete, fut) @@ -822,8 +830,8 @@ class EventLoopTestsMixin: m_socket.error = socket.error m_socket.socket.return_value.setsockopt.side_effect = socket.error - fut = self.event_loop.create_datagram_connection( - protocols.DatagramProtocol) + fut = tasks.Task(self.event_loop.create_datagram_connection( + protocols.DatagramProtocol)) self.assertRaises( socket.error, self.event_loop.run_until_complete, fut) self.assertTrue( diff --git a/tests/tasks_test.py b/tests/tasks_test.py index 2a11c20..03eb526 100644 --- a/tests/tasks_test.py +++ b/tests/tasks_test.py @@ -122,6 +122,39 @@ class TaskTests(test_utils.LogTrackingTestCase): self.assertTrue(t.done()) self.assertFalse(t.cancel()) + def test_future_timeout(self): + @tasks.coroutine + def coro(): + yield from tasks.sleep(10.0) + return 12 + + t = tasks.Task(coro(), timeout=0.1) + + self.assertRaises( + futures.CancelledError, + self.event_loop.run_until_complete, t) + self.assertTrue(t.done()) + self.assertFalse(t.cancel()) + + def test_future_timeout_catch(self): + @tasks.coroutine + def coro(): + yield from tasks.sleep(10.0) + return 12 + + err = None + + @tasks.coroutine + def coro2(): + nonlocal err + try: + yield from tasks.Task(coro(), timeout=0.1) + except futures.CancelledError as exc: + err = exc + + self.event_loop.run_until_complete(tasks.Task(coro2())) + self.assertIsInstance(err, futures.CancelledError) + def test_cancel_in_coro(self): @tasks.coroutine def task(): diff --git a/tulip/base_events.py b/tulip/base_events.py index c2cf18a..b71be09 100644 --- a/tulip/base_events.py +++ b/tulip/base_events.py @@ -239,7 +239,7 @@ class BaseEventLoop(events.AbstractEventLoop): def getnameinfo(self, sockaddr, flags=0): return self.run_in_executor(None, socket.getnameinfo, sockaddr, flags) - @tasks.task + @tasks.coroutine def create_connection(self, protocol_factory, host=None, port=None, *, ssl=False, family=0, proto=0, flags=0, sock=None): """XXX""" @@ -298,7 +298,7 @@ class BaseEventLoop(events.AbstractEventLoop): yield from waiter return transport, protocol - @tasks.task + @tasks.coroutine def create_datagram_connection(self, protocol_factory, host=None, port=None, *, family=socket.AF_INET, proto=0, flags=0): diff --git a/tulip/futures.py b/tulip/futures.py index 68735f3..39137aa 100644 --- a/tulip/futures.py +++ b/tulip/futures.py @@ -51,10 +51,11 @@ class Future: _state = _PENDING _result = None _exception = None + _timeout_handle = None _blocking = False # proper use of future (yield vs yield from) - def __init__(self, *, event_loop=None): + def __init__(self, *, event_loop=None, timeout=None): """Initialize the future. The optional event_loop argument allows to explicitly set the event @@ -67,6 +68,10 @@ class Future: self._event_loop = event_loop self._callbacks = [] + if timeout is not None: + self._timeout_handle = self._event_loop.call_later( + timeout, self.cancel) + def __repr__(self): res = self.__class__.__name__ if self._state == _FINISHED: @@ -105,9 +110,15 @@ class Future: The callbacks are scheduled to be called as soon as possible. Also clears the callback list. """ + # Cancel timeout handle + if self._timeout_handle is not None: + self._timeout_handle.cancel() + self._timeout_handle = None + callbacks = self._callbacks[:] if not callbacks: return + self._callbacks[:] = [] for callback in callbacks: self._event_loop.call_soon(callback, self) diff --git a/tulip/locks.py b/tulip/locks.py index c86048f..4024796 100644 --- a/tulip/locks.py +++ b/tulip/locks.py @@ -94,11 +94,7 @@ class Lock: self._locked = True return True - fut = futures.Future(event_loop=self._event_loop) - if timeout is not None: - handle = self._event_loop.call_later(timeout, fut.cancel) - else: - handle = None + fut = futures.Future(event_loop=self._event_loop, timeout=timeout) self._waiters.append(fut) try: @@ -110,9 +106,6 @@ class Lock: f = self._waiters.popleft() assert f is fut - if handle is not None: - handle.cancel() - self._locked = True return True @@ -209,11 +202,7 @@ class EventWaiter: if self._value: return True - fut = futures.Future(event_loop=self._event_loop) - if timeout is not None: - handle = self._event_loop.call_later(timeout, fut.cancel) - else: - handle = None + fut = futures.Future(event_loop=self._event_loop, timeout=timeout) self._waiters.append(fut) try: @@ -225,9 +214,6 @@ class EventWaiter: f = self._waiters.popleft() assert f is fut - if handle is not None: - handle.cancel() - return True @@ -267,11 +253,7 @@ class Condition(Lock): self.release() - fut = futures.Future(event_loop=self._event_loop) - if timeout is not None: - handle = self._event_loop.call_later(timeout, fut.cancel) - else: - handle = None + fut = futures.Future(event_loop=self._event_loop, timeout=timeout) self._condition_waiters.append(fut) try: @@ -285,9 +267,6 @@ class Condition(Lock): finally: yield from self.acquire() - if handle is not None: - handle.cancel() - return True @tasks.coroutine @@ -406,11 +385,7 @@ class Semaphore: self._locked = True return True - fut = futures.Future(event_loop=self._event_loop) - if timeout is not None: - handle = self._event_loop.call_later(timeout, fut.cancel) - else: - handle = None + fut = futures.Future(event_loop=self._event_loop, timeout=timeout) self._waiters.append(fut) try: @@ -422,9 +397,6 @@ class Semaphore: f = self._waiters.popleft() assert f is fut - if handle is not None: - handle.cancel() - self._value -= 1 if self._value == 0: self._locked = True diff --git a/tulip/tasks.py b/tulip/tasks.py index 08bbb31..45aba01 100644 --- a/tulip/tasks.py +++ b/tulip/tasks.py @@ -10,7 +10,6 @@ import inspect import logging import time -from . import events from . import futures @@ -46,9 +45,9 @@ def task(func): class Task(futures.Future): """A coroutine wrapped in a Future.""" - def __init__(self, coro, event_loop=None): + def __init__(self, coro, event_loop=None, timeout=None): assert inspect.isgenerator(coro) # Must be a coroutine *object*. - super().__init__(event_loop=event_loop) # Sets self._event_loop. + super().__init__(event_loop=event_loop, timeout=timeout) self._coro = coro self._must_cancel = False self._event_loop.call_soon(self._step) @@ -202,13 +201,8 @@ def _wait(fs, timeout=None, return_when=ALL_COMPLETED): return_when == FIRST_EXCEPTION and errors): return done, pending - bail = futures.Future() # Will always be cancelled eventually. - timeout_handle = None - debugstuff = locals() - - if timeout is not None: - loop = events.get_event_loop() - timeout_handle = loop.call_later(timeout, bail.cancel) + # Will always be cancelled eventually. + bail = futures.Future(timeout=timeout) def _on_completion(f): pending.remove(f) @@ -230,8 +224,6 @@ def _wait(fs, timeout=None, return_when=ALL_COMPLETED): finally: for f in pending: f.remove_done_callback(_on_completion) - if timeout_handle is not None: - timeout_handle.cancel() really_done = set(f for f in pending if f.done()) if really_done: |
