diff --git a/scrapy/trunk/scrapy/utils/decompressor.py b/scrapy/trunk/scrapy/contrib/downloadermiddleware/decompression.py similarity index 63% rename from scrapy/trunk/scrapy/utils/decompressor.py rename to scrapy/trunk/scrapy/contrib/downloadermiddleware/decompression.py index 67849ef48..d8ee14ed5 100644 --- a/scrapy/trunk/scrapy/utils/decompressor.py +++ b/scrapy/trunk/scrapy/contrib/downloadermiddleware/decompression.py @@ -1,25 +1,27 @@ -""" -Utility for autodetecting and decompressing responses -""" - import zipfile import tarfile import gzip import bz2 -from scrapy.http import ResponseBody try: from cStringIO import StringIO except: from StringIO import StringIO -class Decompressor(object): - class ArchiveIsEmpty(Exception): - pass - +from scrapy import log +from scrapy.http import Response, ResponseBody + +class DecompressionMiddleware(object): + """ This middleware tries to recognise and extract the possibly compressed + responses that may arrive. """ + def __init__(self): - self.decompressors = {'tar': self.is_tar, 'zip': self.is_zip, - 'gz': self.is_gzip, 'bz2': self.is_bzip2} - + self.decompressors = { + 'tar': self.is_tar, + 'zip': self.is_zip, + 'gz': self.is_gzip, + 'bz2': self.is_bzip2 + } + def is_tar(self, response): try: tar_file = tarfile.open(name='tar.tmp', fileobj=self.archive) @@ -29,7 +31,7 @@ class Decompressor(object): return response.replace(body=ResponseBody(tar_file.extractfile(tar_file.members[0]).read())) else: raise self.ArchiveIsEmpty - + def is_zip(self, response): try: zip_file = zipfile.ZipFile(self.archive) @@ -40,7 +42,7 @@ class Decompressor(object): return response.replace(body=ResponseBody(zip_file.read(namelist[0]))) else: raise self.ArchiveIsEmpty - + def is_gzip(self, response): try: gzip_file = gzip.GzipFile(fileobj=self.archive) @@ -48,25 +50,33 @@ class Decompressor(object): except IOError: return False return response.replace(body=decompressed_body) - + def is_bzip2(self, response): try: decompressed_body = bz2.decompress(self.body) except IOError: return False return response.replace(body=ResponseBody(decompressed_body)) - - def extract_winfo(self, response): + + def extract(self, response): + """ This method tries to decompress the given response, if possible, + and returns a tuple containing the resulting response, and the name + of the used decompressor """ + self.body = response.body.to_string() self.archive = StringIO() self.archive.write(self.body) - + for decompressor in self.decompressors.keys(): self.archive.seek(0) new_response = self.decompressors[decompressor](response) if new_response: - return new_response, decompressor - return response, '' - - def extract(self, response): - return self.extract_winfo(response)[0] + return (new_response, decompressor) + return (response, None) + + def process_response(self, request, response, spider): + if isinstance(response, Response): + response, format = self.extract(response) + if format: + log.msg('Decompressed response with format: %s' % format, log.DEBUG, domain=spider.domain_name) + return response diff --git a/scrapy/trunk/scrapy/contrib/spiders/feed.py b/scrapy/trunk/scrapy/contrib/spiders/feed.py index 502c952b5..37f9dcf6d 100644 --- a/scrapy/trunk/scrapy/contrib/spiders/feed.py +++ b/scrapy/trunk/scrapy/contrib/spiders/feed.py @@ -3,7 +3,6 @@ from scrapy.spider import BaseSpider from scrapy.item import ScrapedItem from scrapy.http import Request from scrapy.utils.iterators import xmliter, csviter -from scrapy.utils.decompressor import Decompressor from scrapy.xpath.selector import XmlXPathSelector from scrapy.core.exceptions import UsageError, NotConfigured @@ -41,9 +40,6 @@ class XMLFeedSpider(BaseSpider): if not hasattr(self, 'parse_item'): raise NotConfigured('You must define parse_item method in order to scrape this XML feed') - decompressor = Decompressor() - response = decompressor.extract(response) - if self.iternodes: nodes = xmliter(response, self.itertag) else: @@ -85,10 +81,6 @@ class CSVFeedSpider(BaseSpider): def parse(self, response): if not hasattr(self, 'parse_row'): raise NotConfigured('You must define parse_row method in order to scrape this CSV feed') - - decompressor = Decompressor() - response = decompressor.extract(response) - response = self.adapt_response(response) return self.parse_rows(response) diff --git a/scrapy/trunk/scrapy/tests/test_decompress.py b/scrapy/trunk/scrapy/tests/test_middleware_decompression.py similarity index 54% rename from scrapy/trunk/scrapy/tests/test_decompress.py rename to scrapy/trunk/scrapy/tests/test_middleware_decompression.py index 784495b76..cc1612c8d 100644 --- a/scrapy/trunk/scrapy/tests/test_decompress.py +++ b/scrapy/trunk/scrapy/tests/test_middleware_decompression.py @@ -1,12 +1,12 @@ import os from unittest import TestCase, main from scrapy.http import Response, ResponseBody -from scrapy.utils.decompressor import Decompressor +from scrapy.contrib.downloadermiddleware.decompression import DecompressionMiddleware -class ScrapyDecompressTest(TestCase): +class ScrapyDecompressionTest(TestCase): uncompressed_body = '' test_responses = {} - decompressor = Decompressor() + middleware = DecompressionMiddleware() def setUp(self): self.datadir = os.path.join(os.path.abspath(os.path.dirname(__file__)), 'sample_data', 'compressed') @@ -22,24 +22,20 @@ class ScrapyDecompressTest(TestCase): self.test_responses[format] = Response('foo.com', 'http://foo.com/bar', body=body) def test_tar(self): - ret = self.decompressor.extract(self.test_responses['tar']) - if ret: - self.assertEqual(ret.body.to_string(), self.uncompressed_body) + response, format = self.middleware.extract(self.test_responses['tar']) + self.assertEqual(response.body.to_string(), self.uncompressed_body) def test_zip(self): - ret = self.decompressor.extract(self.test_responses['zip']) - if ret: - self.assertEqual(ret.body.to_string(), self.uncompressed_body) + response, format = self.middleware.extract(self.test_responses['zip']) + self.assertEqual(response.body.to_string(), self.uncompressed_body) def test_gz(self): - ret = self.decompressor.extract(self.test_responses['xml.gz']) - if ret: - self.assertEqual(ret.body.to_string(), self.uncompressed_body) + response, format = self.middleware.extract(self.test_responses['xml.gz']) + self.assertEqual(response.body.to_string(), self.uncompressed_body) def test_bz2(self): - ret = self.decompressor.extract(self.test_responses['xml.bz2']) - if ret: - self.assertEqual(ret.body.to_string(), self.uncompressed_body) + response, format = self.middleware.extract(self.test_responses['xml.bz2']) + self.assertEqual(response.body.to_string(), self.uncompressed_body) if __name__ == '__main__': main()