mirror of https://github.com/scrapy/scrapy.git
Added a weak key factory based cache
--HG-- extra : rebase_source : 2bc7cb5fdb0fd3adb63cf7fe3aedd2f1d15e49f0
This commit is contained in:
parent
00d55fbbd1
commit
8d1d3493e7
|
|
@ -3,27 +3,26 @@ DefaultHeaders downloader middleware
|
|||
|
||||
See documentation in docs/topics/downloader-middleware.rst
|
||||
"""
|
||||
from scrapy.conf import settings
|
||||
from scrapy.xlib.pydispatch import dispatcher
|
||||
from scrapy import signals
|
||||
from scrapy import conf
|
||||
from scrapy.utils.python import WeakKeyCache
|
||||
|
||||
|
||||
class DefaultHeadersMiddleware(object):
|
||||
|
||||
def __init__(self):
|
||||
def __init__(self, settings=conf.settings):
|
||||
self.global_default_headers = settings.get('DEFAULT_REQUEST_HEADERS')
|
||||
self._default_headers = {}
|
||||
dispatcher.connect(self.spider_opened, signal=signals.spider_opened)
|
||||
dispatcher.connect(self.spider_closed, signal=signals.spider_closed)
|
||||
self._headers = WeakKeyCache(self._default_headers)
|
||||
|
||||
def _default_headers(self, spider):
|
||||
headers = dict(self.global_default_headers)
|
||||
spider_headers = getattr(spider, 'default_request_headers', None) or {}
|
||||
for k, v in spider_headers.iteritems():
|
||||
if v:
|
||||
headers[k] = v
|
||||
else:
|
||||
headers.pop(k, None)
|
||||
return headers.items()
|
||||
|
||||
def process_request(self, request, spider):
|
||||
for k, v in self._default_headers[spider].iteritems():
|
||||
if v:
|
||||
request.headers.setdefault(k, v)
|
||||
|
||||
def spider_opened(self, spider):
|
||||
self._default_headers[spider] = dict(self.global_default_headers,
|
||||
**getattr(spider, 'default_request_headers', {}))
|
||||
|
||||
def spider_closed(self, spider):
|
||||
self._default_headers.pop(spider)
|
||||
for k, v in self._headers[spider]:
|
||||
request.headers.setdefault(k, v)
|
||||
|
|
|
|||
|
|
@ -4,14 +4,24 @@ HTTP basic auth downloader middleware
|
|||
See documentation in docs/topics/downloader-middleware.rst
|
||||
"""
|
||||
|
||||
from scrapy.utils.request import request_authenticate
|
||||
from scrapy.utils.http import basic_auth_header
|
||||
from scrapy.utils.python import WeakKeyCache
|
||||
|
||||
|
||||
class HttpAuthMiddleware(object):
|
||||
"""This middleware allows spiders to use HTTP auth in a cleaner way
|
||||
"""Set Basic HTTP Authorization header
|
||||
(http_user and http_pass spider class attributes)"""
|
||||
|
||||
def __init__(self):
|
||||
self._cache = WeakKeyCache(self._authorization)
|
||||
|
||||
def _authorization(self, spider):
|
||||
usr = getattr(spider, 'http_user', '')
|
||||
pwd = getattr(spider, 'http_pass', '')
|
||||
if usr or pwd:
|
||||
return basic_auth_header(usr, pwd)
|
||||
|
||||
def process_request(self, request, spider):
|
||||
http_user = getattr(spider, 'http_user', '')
|
||||
http_pass = getattr(spider, 'http_pass', '')
|
||||
if http_user or http_pass:
|
||||
request_authenticate(request, http_user, http_pass)
|
||||
auth = self._cache[spider]
|
||||
if auth and 'Authorization' not in request.headers:
|
||||
request.headers['Authorization'] = auth
|
||||
|
|
|
|||
|
|
@ -1,14 +1,20 @@
|
|||
"""Set User-Agent header per spider or use a default value from settings"""
|
||||
|
||||
from scrapy.conf import settings
|
||||
from scrapy.utils.python import WeakKeyCache
|
||||
|
||||
|
||||
class UserAgentMiddleware(object):
|
||||
"""This middleware allows spiders to override the user_agent"""
|
||||
|
||||
default_useragent = settings.get('USER_AGENT')
|
||||
def __init__(self, settings=settings):
|
||||
self.cache = WeakKeyCache(self._user_agent)
|
||||
self.default_useragent = settings.get('USER_AGENT')
|
||||
|
||||
def _user_agent(self, spider):
|
||||
return getattr(spider, 'user_agent', None) or self.default_useragent
|
||||
|
||||
def process_request(self, request, spider):
|
||||
ua = getattr(spider, 'user_agent', None) or self.default_useragent
|
||||
ua = self.cache[spider]
|
||||
if ua:
|
||||
request.headers.setdefault('User-Agent', ua)
|
||||
|
|
|
|||
|
|
@ -16,9 +16,7 @@ class TestDefaultHeadersMiddleware(TestCase):
|
|||
|
||||
def test_process_request(self):
|
||||
req = Request('http://www.scrapytest.org')
|
||||
self.mw.spider_opened(self.spider)
|
||||
self.mw.process_request(req, self.spider)
|
||||
self.mw.spider_closed(self.spider)
|
||||
self.assertEquals(req.headers, self.default_request_headers)
|
||||
|
||||
def test_spider_default_request_headers(self):
|
||||
|
|
@ -30,9 +28,7 @@ class TestDefaultHeadersMiddleware(TestCase):
|
|||
self.spider.default_request_headers = spider_headers
|
||||
|
||||
req = Request('http://www.scrapytest.org')
|
||||
self.mw.spider_opened(self.spider)
|
||||
self.mw.process_request(req, self.spider)
|
||||
self.mw.spider_closed(self.spider)
|
||||
self.assertEquals(req.headers, dict(self.default_request_headers, **spider_headers))
|
||||
|
||||
def test_update_headers(self):
|
||||
|
|
@ -40,9 +36,7 @@ class TestDefaultHeadersMiddleware(TestCase):
|
|||
req = Request('http://www.scrapytest.org', headers=headers)
|
||||
self.assertEquals(req.headers, headers)
|
||||
|
||||
self.mw.spider_opened(self.spider)
|
||||
self.mw.process_request(req, self.spider)
|
||||
self.mw.spider_closed(self.spider)
|
||||
self.default_request_headers.update(headers)
|
||||
self.assertEquals(req.headers, self.default_request_headers)
|
||||
|
||||
|
|
|
|||
|
|
@ -12,17 +12,21 @@ class HttpAuthMiddlewareTest(unittest.TestCase):
|
|||
|
||||
def setUp(self):
|
||||
self.mw = HttpAuthMiddleware()
|
||||
self.spider = TestSpider('foo')
|
||||
|
||||
def tearDown(self):
|
||||
del self.mw
|
||||
|
||||
def test_auth(self):
|
||||
self.mw.default_useragent = 'default_useragent'
|
||||
spider = TestSpider('foo')
|
||||
req = Request('http://scrapytest.org/')
|
||||
assert self.mw.process_request(req, spider) is None
|
||||
assert self.mw.process_request(req, self.spider) is None
|
||||
self.assertEquals(req.headers['Authorization'], 'Basic Zm9vOmJhcg==')
|
||||
|
||||
def test_auth_already_set(self):
|
||||
req = Request('http://scrapytest.org/', headers=dict(Authorization='Digest 123'))
|
||||
assert self.mw.process_request(req, self.spider) is None
|
||||
self.assertEquals(req.headers['Authorization'], 'Digest 123')
|
||||
|
||||
|
||||
if __name__ == '__main__':
|
||||
unittest.main()
|
||||
|
|
|
|||
|
|
@ -0,0 +1,12 @@
|
|||
import unittest
|
||||
from scrapy.utils.http import basic_auth_header
|
||||
|
||||
|
||||
class UtilsHttpTestCase(unittest.TestCase):
|
||||
|
||||
def test_basic_auth_header(self):
|
||||
self.assertEqual('Basic c29tZXVzZXI6c29tZXBhc3M=',
|
||||
basic_auth_header('someuser', 'somepass'))
|
||||
# Check url unsafe encoded header
|
||||
self.assertEqual('Basic c29tZXVzZXI6QDx5dTk-Jm8_UQ==',
|
||||
basic_auth_header('someuser', '@<yu9>&o?Q'))
|
||||
|
|
@ -1,8 +1,10 @@
|
|||
import operator
|
||||
import unittest
|
||||
from itertools import count
|
||||
|
||||
from scrapy.utils.python import str_to_unicode, unicode_to_str, \
|
||||
memoizemethod_noargs, isbinarytext, equal_attributes
|
||||
memoizemethod_noargs, isbinarytext, equal_attributes, \
|
||||
WeakKeyCache
|
||||
|
||||
class UtilsPythonTestCase(unittest.TestCase):
|
||||
def test_str_to_unicode(self):
|
||||
|
|
@ -108,6 +110,19 @@ class UtilsPythonTestCase(unittest.TestCase):
|
|||
a.meta['z'] = 2
|
||||
self.failIf(equal_attributes(a, b, [compare_z, 'x']))
|
||||
|
||||
def test_weakkeycache(self):
|
||||
class _Weakme(object): pass
|
||||
_values = count()
|
||||
wk = WeakKeyCache(lambda k: _values.next())
|
||||
k = _Weakme()
|
||||
v = wk[k]
|
||||
self.assertEqual(v, wk[k])
|
||||
self.assertNotEqual(v, wk[_Weakme()])
|
||||
self.assertEqual(v, wk[k])
|
||||
del k
|
||||
self.assertFalse(len(wk._weakdict))
|
||||
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
unittest.main()
|
||||
|
|
|
|||
|
|
@ -1,3 +1,4 @@
|
|||
from base64 import urlsafe_b64encode
|
||||
|
||||
def headers_raw_to_dict(headers_raw):
|
||||
"""
|
||||
|
|
@ -55,4 +56,6 @@ def headers_dict_to_raw(headers_dict):
|
|||
return '\r\n'.join(raw_lines)
|
||||
|
||||
|
||||
|
||||
def basic_auth_header(username, password):
|
||||
"""Return `Authorization` header for HTTP Basic Access Authentication (RFC 2617)"""
|
||||
return 'Basic ' + urlsafe_b64encode("%s:%s" % (username, password))
|
||||
|
|
|
|||
|
|
@ -11,6 +11,7 @@ import weakref
|
|||
from functools import wraps
|
||||
from sgmllib import SGMLParser
|
||||
|
||||
|
||||
class FixedSGMLParser(SGMLParser):
|
||||
"""The SGMLParser that comes with Python has a bug in the convert_charref()
|
||||
method. This is the same class with the bug fixed"""
|
||||
|
|
@ -179,3 +180,14 @@ def equal_attributes(obj1, obj2, attributes):
|
|||
# all attributes equal
|
||||
return True
|
||||
|
||||
|
||||
class WeakKeyCache(object):
|
||||
|
||||
def __init__(self, default_factory):
|
||||
self.default_factory = default_factory
|
||||
self._weakdict = weakref.WeakKeyDictionary()
|
||||
|
||||
def __getitem__(self, key):
|
||||
if key not in self._weakdict:
|
||||
self._weakdict[key] = self.default_factory(key)
|
||||
return self._weakdict[key]
|
||||
|
|
|
|||
|
|
@ -10,6 +10,7 @@ from urlparse import urlunparse
|
|||
|
||||
from scrapy.utils.url import canonicalize_url
|
||||
from scrapy.utils.httpobj import urlparse_cached
|
||||
from scrapy.utils.http import basic_auth_header
|
||||
|
||||
|
||||
_fingerprint_cache = weakref.WeakKeyDictionary()
|
||||
|
|
@ -61,8 +62,7 @@ def request_authenticate(request, username, password):
|
|||
"""Autenticate the given request (in place) using the HTTP basic access
|
||||
authentication mechanism (RFC 2617) and the given username and password
|
||||
"""
|
||||
b64userpass = urlsafe_b64encode("%s:%s" % (username, password))
|
||||
request.headers['Authorization'] = 'Basic ' + b64userpass
|
||||
request.headers['Authorization'] = basic_auth_header(username, password)
|
||||
|
||||
def request_info(request):
|
||||
"""Return a short string with request info including method, url and
|
||||
|
|
|
|||
Loading…
Reference in New Issue