From 95fde0a4987acaa75a6749223c8b7f9bd7081c23 Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?Daniel=20Gra=C3=B1a?= Date: Tue, 29 Jan 2013 12:00:13 -0200 Subject: [PATCH] register namespaces when XMLFeedSpider uses iternodes mode. fixes #12 --- scrapy/contrib/spiders/feed.py | 20 ++++++++------- scrapy/tests/test_spider.py | 47 +++++++++++++++++++++++++++++++++- 2 files changed, 57 insertions(+), 10 deletions(-) diff --git a/scrapy/contrib/spiders/feed.py b/scrapy/contrib/spiders/feed.py index 10b837353..89c6277ab 100644 --- a/scrapy/contrib/spiders/feed.py +++ b/scrapy/contrib/spiders/feed.py @@ -4,14 +4,15 @@ for scraping from an XML feed. See documentation in docs/topics/spiders.rst """ - from scrapy.spider import BaseSpider from scrapy.item import BaseItem from scrapy.http import Request from scrapy.utils.iterators import xmliter, csviter +from scrapy.utils.spider import iterate_spider_output from scrapy.selector import XmlXPathSelector, HtmlXPathSelector from scrapy.exceptions import NotConfigured, NotSupported + class XMLFeedSpider(BaseSpider): """ This class intends to be the base class for spiders that scrape @@ -45,10 +46,10 @@ class XMLFeedSpider(BaseSpider): def parse_node(self, response, selector): """This method must be overriden with your custom spider functionality""" - if hasattr(self, 'parse_item'): # backward compatibility + if hasattr(self, 'parse_item'): # backward compatibility return self.parse_item(response, selector) raise NotImplementedError - + def parse_nodes(self, response, nodes): """This method is called for the nodes matching the provided tag name (itertag). Receives the response and an XPathSelector for each node. @@ -58,11 +59,7 @@ class XMLFeedSpider(BaseSpider): """ for selector in nodes: - ret = self.parse_node(response, selector) - if isinstance(ret, (BaseItem, Request)): - ret = [ret] - if not isinstance(ret, (list, tuple)): - raise TypeError('You cannot return an "%s" object from a spider' % type(ret).__name__) + ret = iterate_spider_output(self.parse_node(response, selector)) for result_item in self.process_results(response, ret): yield result_item @@ -72,7 +69,7 @@ class XMLFeedSpider(BaseSpider): response = self.adapt_response(response) if self.iterator == 'iternodes': - nodes = xmliter(response, self.itertag) + nodes = self._iternodes(response) elif self.iterator == 'xml': selector = XmlXPathSelector(response) self._register_namespaces(selector) @@ -86,6 +83,11 @@ class XMLFeedSpider(BaseSpider): return self.parse_nodes(response, nodes) + def _iternodes(self, response): + for node in xmliter(response, self.itertag): + self._register_namespaces(node) + yield node + def _register_namespaces(self, selector): for (prefix, uri) in self.namespaces: selector.register_namespace(prefix, uri) diff --git a/scrapy/tests/test_spider.py b/scrapy/tests/test_spider.py index bca8de4b9..9ffa2d7cc 100644 --- a/scrapy/tests/test_spider.py +++ b/scrapy/tests/test_spider.py @@ -1,4 +1,6 @@ -import gzip, warnings, inspect +import gzip +import inspect +import warnings from cStringIO import StringIO from twisted.trial import unittest @@ -45,18 +47,61 @@ class InitSpiderTest(BaseSpiderTest): spider_class = InitSpider + class XMLFeedSpiderTest(BaseSpiderTest): spider_class = XMLFeedSpider + def test_register_namespace(self): + body = """ + + http://www.example.com/Special-Offers.html2009-08-16 + http://www.example.com/2009-08-16 + """ + response = XmlResponse(url='http://example.com/sitemap.xml', body=body) + + class _XMLSpider(self.spider_class): + itertag = 'url' + namespaces = ( + ('a', 'http://www.google.com/schemas/sitemap/0.84'), + ('b', 'http://www.example.com/schemas/extras/1.0'), + ) + + def parse_node(self, response, selector): + yield { + 'loc': selector.select('a:loc/text()').extract(), + 'updated': selector.select('b:updated/text()').extract(), + 'other': selector.select('other/@value').extract(), + 'custom': selector.select('other/@b:custom').extract(), + } + + for iterator in ('iternodes', 'xml'): + spider = _XMLSpider('example', iterator=iterator) + output = list(spider.parse(response)) + self.assertEqual(len(output), 2, iterator) + self.assertEqual(output, [ + {'loc': [u'http://www.example.com/Special-Offers.html'], + 'updated': [u'2009-08-16'], + 'custom': [u'fuu'], + 'other': [u'bar']}, + {'loc': [], + 'updated': [u'2009-08-16'], + 'other': [u'foo'], + 'custom': []}, + ], iterator) + + class CSVFeedSpiderTest(BaseSpiderTest): spider_class = CSVFeedSpider + class CrawlSpiderTest(BaseSpiderTest): spider_class = CrawlSpider + class SitemapSpiderTest(BaseSpiderTest): spider_class = SitemapSpider