From 2822db730fc986ff150e88ce714d8abcfc198020 Mon Sep 17 00:00:00 2001 From: Pablo Hoffman Date: Tue, 10 Aug 2010 18:03:47 -0300 Subject: [PATCH] added new MiddlewareManager class that will be used as base class for pipeline and middlewares --- scrapy/middleware.py | 63 ++++++++++++++++++++++++ scrapy/tests/test_middleware.py | 84 ++++++++++++++++++++++++++++++++ scrapy/tests/test_utils_defer.py | 55 +++++++++++++++++++-- scrapy/utils/defer.py | 29 +++++++++++ 4 files changed, 226 insertions(+), 5 deletions(-) create mode 100644 scrapy/middleware.py create mode 100644 scrapy/tests/test_middleware.py diff --git a/scrapy/middleware.py b/scrapy/middleware.py new file mode 100644 index 000000000..148b9fd50 --- /dev/null +++ b/scrapy/middleware.py @@ -0,0 +1,63 @@ +from collections import defaultdict + +from scrapy import log +from scrapy.exceptions import NotConfigured +from scrapy.utils.misc import load_object +from scrapy.utils.defer import process_parallel, process_chain, process_chain_both + +class MiddlewareManager(object): + """Base class for implementing middleware managers""" + + component_name = 'foo middleware' + + def __init__(self, *middlewares): + self.middlewares = middlewares + self.methods = defaultdict(list) + for mw in middlewares: + self._add_middleware(mw) + + @classmethod + def _get_mwlist_from_settings(cls, settings): + return [] + + @classmethod + def from_settings(cls, settings): + mwlist = cls._get_mwlist_from_settings(settings) + middlewares = [] + for clspath in mwlist: + try: + mwcls = load_object(clspath) + if hasattr(mwcls, 'from_settings'): + mw = mwcls.from_settings(settings) + else: + mw = mwcls() + middlewares.append(mw) + except NotConfigured, e: + if e.args: + log.msg(e) + enabled = [type(x).__name__ for x in middlewares] + log.msg("Enabled %ss: %s" % (cls.component_name, ",".join(enabled)), \ + level=log.DEBUG) + return cls(*middlewares) + + def _add_middleware(self, mw): + if hasattr(mw, 'open_spider'): + self.methods['open_spider'].append(mw.open_spider) + if hasattr(mw, 'close_spider'): + self.methods['close_spider'].insert(0, mw.close_spider) + + def _process_parallel(self, methodname, obj, *args): + return process_parallel(self.methods[methodname], obj, *args) + + def _process_chain(self, methodname, obj, *args): + return process_chain(self.methods[methodname], obj, *args) + + def _process_chain_both(self, cb_methodname, eb_methodname, obj, *args): + return process_chain_both(self.methods[cb_methodname], \ + self.methods[eb_methodname], obj, *args) + + def open_spider(self, spider): + return self._process_parallel('open_spider', spider) + + def close_spider(self, spider): + return self._process_parallel('close_spider', spider) diff --git a/scrapy/tests/test_middleware.py b/scrapy/tests/test_middleware.py new file mode 100644 index 000000000..1eca89b27 --- /dev/null +++ b/scrapy/tests/test_middleware.py @@ -0,0 +1,84 @@ +from twisted.trial import unittest + +from scrapy.conf import Settings +from scrapy.exceptions import NotConfigured +from scrapy.middleware import MiddlewareManager + +class M1(object): + + def open_spider(self, spider): + pass + + def close_spider(self, spider): + pass + + def process(self, response, request, spider): + pass + +class M2(object): + + def open_spider(self, spider): + pass + + def close_spider(self, spider): + pass + + pass + +class M3(object): + + def process(self, response, request, spider): + pass + + +class MOff(object): + + def open_spider(self, spider): + pass + + def close_spider(self, spider): + pass + + def __init__(self): + raise NotConfigured + + +class TestMiddlewareManager(MiddlewareManager): + + @classmethod + def _get_mwlist_from_settings(cls, settings): + return ['scrapy.tests.test_middleware.%s' % x for x in ['M1', 'MOff', 'M3']] + + def _add_middleware(self, mw): + super(TestMiddlewareManager, self)._add_middleware(mw) + if hasattr(mw, 'process'): + self.methods['process'].append(mw.process) + +class MiddlewareManagerTest(unittest.TestCase): + + def test_init(self): + m1, m2, m3 = M1(), M2(), M3() + mwman = TestMiddlewareManager(m1, m2, m3) + self.assertEqual(mwman.methods['open_spider'], [m1.open_spider, m2.open_spider]) + self.assertEqual(mwman.methods['close_spider'], [m2.close_spider, m1.close_spider]) + self.assertEqual(mwman.methods['process'], [m1.process, m3.process]) + + def test_methods(self): + mwman = TestMiddlewareManager(M1(), M2(), M3()) + self.assertEqual([x.im_class for x in mwman.methods['open_spider']], + [M1, M2]) + self.assertEqual([x.im_class for x in mwman.methods['close_spider']], + [M2, M1]) + self.assertEqual([x.im_class for x in mwman.methods['process']], + [M1, M3]) + + def test_enabled(self): + m1, m2, m3 = M1(), M2(), M3() + mwman = MiddlewareManager(m1, m2, m3) + self.failUnlessEqual(mwman.middlewares, (m1, m2, m3)) + + def test_enabled_from_settings(self): + settings = Settings() + mwman = TestMiddlewareManager.from_settings(settings) + classes = [x.__class__ for x in mwman.middlewares] + self.failUnlessEqual(classes, [M1, M3]) diff --git a/scrapy/tests/test_utils_defer.py b/scrapy/tests/test_utils_defer.py index 84f39c25d..4e140b3f7 100644 --- a/scrapy/tests/test_utils_defer.py +++ b/scrapy/tests/test_utils_defer.py @@ -1,9 +1,9 @@ -from itertools import imap from twisted.trial import unittest +from twisted.internet import reactor, defer +from twisted.python.failure import Failure -from twisted.internet import reactor -from twisted.internet.defer import Deferred -from scrapy.utils.defer import mustbe_deferred +from scrapy.utils.defer import mustbe_deferred, process_chain, \ + process_chain_both, process_parallel class MustbeDeferredTest(unittest.TestCase): @@ -22,7 +22,7 @@ class MustbeDeferredTest(unittest.TestCase): steps = [] def _append(v): steps.append(v) - dfd = Deferred() + dfd = defer.Deferred() reactor.callLater(0, dfd.callback, steps) return dfd @@ -31,3 +31,48 @@ class MustbeDeferredTest(unittest.TestCase): steps.append(2) # add another value, that should be catched by assertEqual return dfd +def cb1(value, arg1, arg2): + return "(cb1 %s %s %s)" % (value, arg1, arg2) +def cb2(value, arg1, arg2): + return defer.succeed("(cb2 %s %s %s)" % (value, arg1, arg2)) +def cb3(value, arg1, arg2): + return "(cb3 %s %s %s)" % (value, arg1, arg2) +def cb_fail(value, arg1, arg2): + return Failure(TypeError()) +def eb1(failure, arg1, arg2): + return "(eb1 %s %s %s)" % (failure.value.__class__.__name__, arg1, arg2) + + +class DeferUtilsTest(unittest.TestCase): + + @defer.inlineCallbacks + def test_process_chain(self): + x = yield process_chain([cb1, cb2, cb3], 'res', 'v1', 'v2') + self.assertEqual(x, "(cb3 (cb2 (cb1 res v1 v2) v1 v2) v1 v2)") + + gotexc = False + try: + yield process_chain([cb1, cb_fail, cb3], 'res', 'v1', 'v2') + except TypeError, e: + gotexc = True + self.failUnless(gotexc) + + @defer.inlineCallbacks + def test_process_chain_both(self): + x = yield process_chain_both([cb_fail, cb2, cb3], [None, eb1, None], 'res', 'v1', 'v2') + self.assertEqual(x, "(cb3 (eb1 TypeError v1 v2) v1 v2)") + + fail = Failure(ZeroDivisionError()) + x = yield process_chain_both([eb1, cb2, cb3], [eb1, None, None], fail, 'v1', 'v2') + self.assertEqual(x, "(cb3 (cb2 (eb1 ZeroDivisionError v1 v2) v1 v2) v1 v2)") + + @defer.inlineCallbacks + def test_process_parallel(self): + x = yield process_parallel([cb1, cb2, cb3], 'res', 'v1', 'v2') + self.assertEqual(x, ['(cb1 res v1 v2)', '(cb2 res v1 v2)', '(cb3 res v1 v2)']) + + def test_process_parallel_failure(self): + d = process_parallel([cb1, cb_fail, cb3], 'res', 'v1', 'v2') + self.failUnlessFailure(d, TypeError) + self.flushLoggedErrors() + return d diff --git a/scrapy/utils/defer.py b/scrapy/utils/defer.py index 7fb34efd5..59abb5bfa 100644 --- a/scrapy/utils/defer.py +++ b/scrapy/utils/defer.py @@ -56,3 +56,32 @@ def parallel(iterable, count, callable, *args, **named): coop = task.Cooperator() work = (callable(elem, *args, **named) for elem in iterable) return defer.DeferredList([coop.coiterate(work) for i in xrange(count)]) + +def process_chain(callbacks, input, *a, **kw): + """Return a Deferred built by chaining the given callbacks""" + d = defer.Deferred() + for x in callbacks: + d.addCallback(x, *a, **kw) + d.callback(input) + return d + +def process_chain_both(callbacks, errbacks, input, *a, **kw): + """Return a Deferred built by chaining the given callbacks and errbacks""" + d = defer.Deferred() + for cb, eb in zip(callbacks, errbacks): + d.addCallbacks(cb, eb, callbackArgs=a, callbackKeywords=kw, + errbackArgs=a, errbackKeywords=kw) + if isinstance(input, failure.Failure): + d.errback(input) + else: + d.callback(input) + return d + +def process_parallel(callbacks, input, *a, **kw): + """Return a Deferred with the output of all successful calls to the given + callbacks + """ + dfds = [defer.succeed(input).addCallback(x, *a, **kw) for x in callbacks] + d = defer.gatherResults(dfds) + d.addErrback(lambda _: _.value.subFailure) + return d