Added support for passing generic arguments to spider constructors (refs #152), extended Spider tests, added unittests for TwistedPluginSpiderManager

This commit is contained in:
Pablo Hoffman 2010-04-05 11:27:19 -03:00
parent de32612c99
commit c99e1af766
8 changed files with 117 additions and 29 deletions

View File

@ -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):

View File

@ -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):

View File

@ -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

View File

@ -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

View File

@ -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()

View File

@ -0,0 +1,7 @@
from scrapy.spider import BaseSpider
class Spider1(BaseSpider):
name = "spider1"
allowed_domains = ["scrapy1.org", "scrapy3.org"]
SPIDER = Spider1()

View File

@ -0,0 +1,7 @@
from scrapy.spider import BaseSpider
class Spider2(BaseSpider):
name = "spider2"
allowed_domains = ["scrapy2.org", "scrapy3.org"]
SPIDER = Spider2()

View File

@ -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()