mirror of https://github.com/scrapy/scrapy.git
added new MiddlewareManager class that will be used as base class for pipeline and middlewares
This commit is contained in:
parent
9d38a99aa8
commit
2822db730f
|
|
@ -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)
|
||||
|
|
@ -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])
|
||||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
Loading…
Reference in New Issue