summaryrefslogtreecommitdiff
diff options
context:
space:
mode:
authorNikolay Kim <fafhrd91@gmail.com>2013-03-20 11:25:20 -0700
committerNikolay Kim <fafhrd91@gmail.com>2013-03-20 11:25:20 -0700
commit35981988fbf103d88d6415be790dfd2aa140260d (patch)
tree322ce1a7bbb54544576429688b7dd713d230dd1b
parentb780325d12b2b0faca58d87792aa26f1067963cc (diff)
downloadtrollius-git-35981988fbf103d88d6415be790dfd2aa140260d.tar.gz
timeout support for Future
-rw-r--r--tests/base_events_test.py4
-rw-r--r--tests/events_test.py48
-rw-r--r--tests/tasks_test.py33
-rw-r--r--tulip/base_events.py4
-rw-r--r--tulip/futures.py13
-rw-r--r--tulip/locks.py36
-rw-r--r--tulip/tasks.py16
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: