mirror of https://github.com/scrapy/scrapy.git
Added support for passing generic arguments to spider constructors (refs #152), extended Spider tests, added unittests for TwistedPluginSpiderManager
This commit is contained in:
parent
de32612c99
commit
c99e1af766
|
|
@ -4,7 +4,6 @@ spiders
|
|||
"""
|
||||
|
||||
import sys
|
||||
import urlparse
|
||||
|
||||
from twisted.plugin import getCache
|
||||
from twisted.python.rebuild import rebuild
|
||||
|
|
@ -21,21 +20,18 @@ class TwistedPluginSpiderManager(object):
|
|||
self.loaded = False
|
||||
self._spiders = {}
|
||||
|
||||
def create(self, spider_id):
|
||||
def create(self, spider_name, **spider_kwargs):
|
||||
"""Returns a Spider instance for the given spider name, using the given
|
||||
spider arguments. If the sipder name is not found, it raises a
|
||||
KeyError.
|
||||
"""
|
||||
Returns Spider instance by given identifier.
|
||||
If not exists raises KeyError.
|
||||
"""
|
||||
#@@@ currently spider_id = domain
|
||||
# if lookup fails let dict's KeyError exception propagate
|
||||
return self._spiders[spider_id]
|
||||
spider = self._spiders[spider_name]
|
||||
spider.__dict__.update(spider_kwargs)
|
||||
return spider
|
||||
|
||||
def find_by_request(self, request):
|
||||
"""
|
||||
Returns list of spiders ids that match given Request.
|
||||
"""
|
||||
# just find by request.url
|
||||
return [domain for domain, spider in self._spiders.iteritems()
|
||||
"""Returns list of spiders names that match the given Request"""
|
||||
return [name for name, spider in self._spiders.iteritems()
|
||||
if url_is_from_spider(request.url, spider)]
|
||||
|
||||
def list(self):
|
||||
|
|
|
|||
|
|
@ -59,9 +59,9 @@ class CrawlSpider(InitSpider):
|
|||
"""
|
||||
rules = ()
|
||||
|
||||
def __init__(self):
|
||||
def __init__(self, *a, **kw):
|
||||
"""Constructor takes care of compiling rules"""
|
||||
super(CrawlSpider, self).__init__()
|
||||
super(CrawlSpider, self).__init__(*a, **kw)
|
||||
self._compile_rules()
|
||||
|
||||
def parse(self, response):
|
||||
|
|
|
|||
|
|
@ -3,8 +3,8 @@ from scrapy.spider import BaseSpider
|
|||
class InitSpider(BaseSpider):
|
||||
"""Base Spider with initialization facilities"""
|
||||
|
||||
def __init__(self):
|
||||
super(InitSpider, self).__init__()
|
||||
def __init__(self, *a, **kw):
|
||||
super(InitSpider, self).__init__(*a, **kw)
|
||||
self._postinit_reqs = []
|
||||
self._init_complete = False
|
||||
self._init_started = False
|
||||
|
|
|
|||
|
|
@ -30,7 +30,8 @@ class BaseSpider(object_ref):
|
|||
start_urls = []
|
||||
allowed_domains = []
|
||||
|
||||
def __init__(self, name=None):
|
||||
def __init__(self, name=None, **kwargs):
|
||||
self.__dict__.update(kwargs)
|
||||
# XXX: SEP-12 backward compatibility (remove for 0.10)
|
||||
if hasattr(self, 'domain_name'):
|
||||
warnings.warn("Spider.domain_name attribute is deprecated, use Spider.name instead", \
|
||||
|
|
@ -53,6 +54,8 @@ class BaseSpider(object_ref):
|
|||
self.domain_name = self.name
|
||||
if not getattr(self, 'extra_domain_names', None):
|
||||
self.extra_domain_names = self.allowed_domains
|
||||
if not self.name:
|
||||
raise ValueError("%s must have a name" % type(self).__name__)
|
||||
|
||||
def log(self, message, level=log.DEBUG):
|
||||
"""Log the given messages at the given log level. Always use this
|
||||
|
|
|
|||
|
|
@ -0,0 +1,39 @@
|
|||
import unittest
|
||||
|
||||
# just a hack to avoid cyclic imports of scrapy.spider when running this test
|
||||
# alone
|
||||
import scrapy.spider
|
||||
from scrapy.contrib.spidermanager import TwistedPluginSpiderManager
|
||||
from scrapy.http import Request
|
||||
|
||||
class TwistedPluginSpiderManagerTest(unittest.TestCase):
|
||||
|
||||
def setUp(self):
|
||||
self.spiderman = TwistedPluginSpiderManager()
|
||||
assert not self.spiderman.loaded
|
||||
self.spiderman.load(['scrapy.tests.test_contrib_spidermanager'])
|
||||
assert self.spiderman.loaded
|
||||
|
||||
def test_list(self):
|
||||
self.assertEqual(set(self.spiderman.list()),
|
||||
set(['spider1', 'spider2']))
|
||||
|
||||
def test_create(self):
|
||||
spider1 = self.spiderman.create("spider1")
|
||||
self.assertEqual(spider1.__class__.__name__, 'Spider1')
|
||||
spider2 = self.spiderman.create("spider2", foo="bar")
|
||||
self.assertEqual(spider2.__class__.__name__, 'Spider2')
|
||||
self.assertEqual(spider2.foo, 'bar')
|
||||
|
||||
def test_find_by_request(self):
|
||||
self.assertEqual(self.spiderman.find_by_request(Request('http://scrapy1.org/test')),
|
||||
['spider1'])
|
||||
self.assertEqual(self.spiderman.find_by_request(Request('http://scrapy2.org/test')),
|
||||
['spider2'])
|
||||
self.assertEqual(set(self.spiderman.find_by_request(Request('http://scrapy3.org/test'))),
|
||||
set(['spider1', 'spider2']))
|
||||
self.assertEqual(self.spiderman.find_by_request(Request('http://scrapy999.org/test')),
|
||||
[])
|
||||
|
||||
if __name__ == '__main__':
|
||||
unittest.main()
|
||||
|
|
@ -0,0 +1,7 @@
|
|||
from scrapy.spider import BaseSpider
|
||||
|
||||
class Spider1(BaseSpider):
|
||||
name = "spider1"
|
||||
allowed_domains = ["scrapy1.org", "scrapy3.org"]
|
||||
|
||||
SPIDER = Spider1()
|
||||
|
|
@ -0,0 +1,7 @@
|
|||
from scrapy.spider import BaseSpider
|
||||
|
||||
class Spider2(BaseSpider):
|
||||
name = "spider2"
|
||||
allowed_domains = ["scrapy2.org", "scrapy3.org"]
|
||||
|
||||
SPIDER = Spider2()
|
||||
|
|
@ -6,20 +6,25 @@ import warnings
|
|||
from twisted.trial import unittest
|
||||
|
||||
from scrapy.spider import BaseSpider
|
||||
from scrapy.contrib.spiders.init import InitSpider
|
||||
from scrapy.contrib.spiders.crawl import CrawlSpider
|
||||
from scrapy.contrib.spiders.feed import XMLFeedSpider, CSVFeedSpider
|
||||
from scrapy.contrib.dupefilter import RequestFingerprintDupeFilter, NullDupeFilter
|
||||
|
||||
class OldSpider(BaseSpider):
|
||||
class BaseSpiderTest(unittest.TestCase):
|
||||
|
||||
domain_name = 'example.com'
|
||||
extra_domain_names = ('example.org', 'example.net')
|
||||
spider_class = BaseSpider
|
||||
|
||||
class NewSpider(BaseSpider):
|
||||
class OldSpider(spider_class):
|
||||
|
||||
name = 'example.com'
|
||||
allowed_domains = ('example.org', 'example.net')
|
||||
domain_name = 'example.com'
|
||||
extra_domain_names = ('example.org', 'example.net')
|
||||
|
||||
class NewSpider(spider_class):
|
||||
|
||||
name = 'example.com'
|
||||
allowed_domains = ('example.org', 'example.net')
|
||||
|
||||
class SpiderTest(unittest.TestCase):
|
||||
|
||||
def setUp(self):
|
||||
warnings.simplefilter("always")
|
||||
|
|
@ -32,25 +37,56 @@ class SpiderTest(unittest.TestCase):
|
|||
# warnings.catch_warnings() was added in Python 2.6
|
||||
raise unittest.SkipTest("This test requires Python 2.6+")
|
||||
with warnings.catch_warnings(record=True) as w:
|
||||
spider = OldSpider()
|
||||
spider = self.OldSpider()
|
||||
self.assertEqual(len(w), 2) # one for domain_name & one for extra_domain_names
|
||||
self.assert_(issubclass(w[-1].category, DeprecationWarning))
|
||||
|
||||
def test_sep12_backwards_compatibility(self):
|
||||
spider = OldSpider()
|
||||
spider = self.OldSpider()
|
||||
self.assertEqual(spider.name, 'example.com')
|
||||
self.assert_('example.com' in spider.allowed_domains, spider.allowed_domains)
|
||||
self.assert_('example.org' in spider.allowed_domains, spider.allowed_domains)
|
||||
self.assert_('example.net' in spider.allowed_domains, spider.allowed_domains)
|
||||
|
||||
spider = NewSpider()
|
||||
spider = self.NewSpider()
|
||||
self.assertEqual(spider.domain_name, 'example.com')
|
||||
self.assert_('example.org' in spider.extra_domain_names, spider.extra_domain_names)
|
||||
self.assert_('example.net' in spider.extra_domain_names, spider.extra_domain_names)
|
||||
|
||||
def test_base_spider(self):
|
||||
spider = BaseSpider("example.com")
|
||||
spider = self.spider_class("example.com")
|
||||
self.assertEqual(spider.name, 'example.com')
|
||||
self.assertEqual(spider.start_urls, [])
|
||||
self.assertEqual(spider.allowed_domains, [])
|
||||
|
||||
def test_spider_args(self):
|
||||
"""Constructor arguments are assigned to spider attributes"""
|
||||
spider = self.spider_class('example.com', foo='bar')
|
||||
self.assertEqual(spider.foo, 'bar')
|
||||
|
||||
def test_spider_without_name(self):
|
||||
"""Constructor arguments are assigned to spider attributes"""
|
||||
spider = self.spider_class('example.com')
|
||||
self.assertRaises(ValueError, self.spider_class)
|
||||
self.assertRaises(ValueError, self.spider_class, somearg='foo')
|
||||
|
||||
|
||||
class InitSpiderTest(BaseSpiderTest):
|
||||
|
||||
spider_class = InitSpider
|
||||
|
||||
class XMLFeedSpiderTest(BaseSpiderTest):
|
||||
|
||||
spider_class = XMLFeedSpider
|
||||
|
||||
class CSVFeedSpiderTest(BaseSpiderTest):
|
||||
|
||||
spider_class = CSVFeedSpider
|
||||
|
||||
class CrawlSpiderTest(BaseSpiderTest):
|
||||
|
||||
spider_class = CrawlSpider
|
||||
|
||||
|
||||
if __name__ == '__main__':
|
||||
unittest.main()
|
||||
|
|
|
|||
Loading…
Reference in New Issue