From c99e1af766bcf463c1bd1b6dcfd5bb3a5e8e9d39 Mon Sep 17 00:00:00 2001 From: Pablo Hoffman Date: Mon, 5 Apr 2010 11:27:19 -0300 Subject: [PATCH] Added support for passing generic arguments to spider constructors (refs #152), extended Spider tests, added unittests for TwistedPluginSpiderManager --- scrapy/contrib/spidermanager.py | 22 +++---- scrapy/contrib/spiders/crawl.py | 4 +- scrapy/contrib/spiders/init.py | 4 +- scrapy/spider/models.py | 5 +- .../test_contrib_spidermanager/__init__.py | 39 +++++++++++++ .../test_contrib_spidermanager/spider1.py | 7 +++ .../test_contrib_spidermanager/spider2.py | 7 +++ scrapy/tests/test_spider.py | 58 +++++++++++++++---- 8 files changed, 117 insertions(+), 29 deletions(-) create mode 100644 scrapy/tests/test_contrib_spidermanager/__init__.py create mode 100644 scrapy/tests/test_contrib_spidermanager/spider1.py create mode 100644 scrapy/tests/test_contrib_spidermanager/spider2.py diff --git a/scrapy/contrib/spidermanager.py b/scrapy/contrib/spidermanager.py index 331f55aa5..1c41b625d 100644 --- a/scrapy/contrib/spidermanager.py +++ b/scrapy/contrib/spidermanager.py @@ -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): diff --git a/scrapy/contrib/spiders/crawl.py b/scrapy/contrib/spiders/crawl.py index 648765b06..2e74dab4e 100644 --- a/scrapy/contrib/spiders/crawl.py +++ b/scrapy/contrib/spiders/crawl.py @@ -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): diff --git a/scrapy/contrib/spiders/init.py b/scrapy/contrib/spiders/init.py index b37591ca5..f759fd81f 100644 --- a/scrapy/contrib/spiders/init.py +++ b/scrapy/contrib/spiders/init.py @@ -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 diff --git a/scrapy/spider/models.py b/scrapy/spider/models.py index 757aaff55..3adf7a470 100644 --- a/scrapy/spider/models.py +++ b/scrapy/spider/models.py @@ -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 diff --git a/scrapy/tests/test_contrib_spidermanager/__init__.py b/scrapy/tests/test_contrib_spidermanager/__init__.py new file mode 100644 index 000000000..d27c32263 --- /dev/null +++ b/scrapy/tests/test_contrib_spidermanager/__init__.py @@ -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() diff --git a/scrapy/tests/test_contrib_spidermanager/spider1.py b/scrapy/tests/test_contrib_spidermanager/spider1.py new file mode 100644 index 000000000..0a9b60989 --- /dev/null +++ b/scrapy/tests/test_contrib_spidermanager/spider1.py @@ -0,0 +1,7 @@ +from scrapy.spider import BaseSpider + +class Spider1(BaseSpider): + name = "spider1" + allowed_domains = ["scrapy1.org", "scrapy3.org"] + +SPIDER = Spider1() diff --git a/scrapy/tests/test_contrib_spidermanager/spider2.py b/scrapy/tests/test_contrib_spidermanager/spider2.py new file mode 100644 index 000000000..52023277f --- /dev/null +++ b/scrapy/tests/test_contrib_spidermanager/spider2.py @@ -0,0 +1,7 @@ +from scrapy.spider import BaseSpider + +class Spider2(BaseSpider): + name = "spider2" + allowed_domains = ["scrapy2.org", "scrapy3.org"] + +SPIDER = Spider2() diff --git a/scrapy/tests/test_spider.py b/scrapy/tests/test_spider.py index 507de4abb..c488dad3b 100644 --- a/scrapy/tests/test_spider.py +++ b/scrapy/tests/test_spider.py @@ -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()