cookies: final new cookie middleware integration into core and enabled by default. refs #73

--HG--
extra : convert_revision : svn%3Ab85faa78-f9eb-468e-a121-7cced6da292c%401081
This commit is contained in:
Daniel Grana 2009-04-22 17:21:46 +00:00
parent 3de3c2a331
commit 1ecf874a75
12 changed files with 190 additions and 157 deletions

View File

@ -61,12 +61,12 @@ DOWNLOADER_MIDDLEWARES = [
# Engine side
'scrapy.contrib.downloadermiddleware.robotstxt.RobotsTxtMiddleware',
'scrapy.contrib.downloadermiddleware.errorpages.ErrorPagesMiddleware',
'scrapy.contrib.downloadermiddleware.cookies.CookiesMiddleware',
'scrapy.contrib.downloadermiddleware.httpauth.HttpAuthMiddleware',
'scrapy.contrib.downloadermiddleware.useragent.UserAgentMiddleware',
'scrapy.contrib.downloadermiddleware.retry.RetryMiddleware',
'scrapy.contrib.downloadermiddleware.common.CommonMiddleware',
'scrapy.contrib.downloadermiddleware.redirect.RedirectMiddleware',
'scrapy.contrib.downloadermiddleware.cookies.CookiesMiddleware',
'scrapy.contrib.downloadermiddleware.httpcompression.HttpCompressionMiddleware',
'scrapy.contrib.downloadermiddleware.debug.CrawlDebug',
'scrapy.contrib.downloadermiddleware.stats.DownloaderStats',

View File

@ -1,27 +1,81 @@
import operator
from itertools import groupby
from collections import defaultdict
from pydispatch import dispatcher
from scrapy.core import signals
from scrapy.utils.misc import dict_updatedefault
from scrapy.http import Response
from scrapy.utils.cookies import CookieJar
from scrapy.core.exceptions import HttpException
from scrapy.conf import settings
from scrapy import log
class CookiesMiddleware(object):
"""This middleware enables working with sites that need cookies"""
debug = settings.getbool('COOKIES_DEBUG')
def __init__(self):
self.cookies = {}
dispatcher.connect(self.domain_open, signals.domain_open)
self.jars = defaultdict(CookieJar)
dispatcher.connect(self.domain_closed, signals.domain_closed)
def process_request(self, request, spider):
if not request.meta.get('dont_merge_cookies', False):
dict_updatedefault(request.cookies, self.cookies[spider.domain_name])
if request.meta.get('dont_merge_cookies', False):
return
jar = self.jars[spider.domain_name]
cookies = self._get_request_cookies(jar, request)
for cookie in cookies:
jar.set_cookie_if_ok(cookie, request)
# set Cookie header
request.headers.pop('Cookie', None)
jar.add_cookie_header(request)
self._debug_cookie(request)
def process_response(self, request, response, spider):
cookies = self.cookies[spider.domain_name]
cookies.update(request.cookies)
if request.meta.get('dont_merge_cookies', False):
return
# extract cookies from Set-Cookie and drop invalid/expired cookies
jar = self.jars[spider.domain_name]
jar.extract_cookies(response, request)
self._debug_set_cookie(response)
# TODO: set current cookies in jar to response.cookies?
return response
def domain_open(self, domain):
self.cookies[domain] = {}
# cookies should be set on non-200 responses too
def process_exception(self, request, exception, spider):
if isinstance(exception, HttpException):
self.process_response(request, exception.response, spider)
def domain_closed(self, domain):
del self.cookies[domain]
self.jars.pop(domain, None)
def _debug_cookie(self, request):
"""log Cookie header for request"""
if self.debug:
c = request.headers.get('Cookie')
c = c and [p.split('=')[0] for p in c.split(';')]
log.msg('Cookie: %s for %s' % (c, request.url), level=log.DEBUG)
def _debug_set_cookie(self, response):
"""log Set-Cookies headers but exclude cookie values"""
if self.debug:
cl = response.headers.getlist('Set-Cookie')
res = []
for c in cl:
kv, tail = c.split(';', 1)
k = kv.split('=', 1)[0]
res.append('%s %s' % (k, tail))
log.msg('Set-Cookie: %s from %s' % (res, response.url))
def _get_request_cookies(self, jar, request):
headers = {'Set-Cookie': ['%s=%s;' % (k, v) for k, v in request.cookies.iteritems()]}
response = Response(request.url, headers=headers)
cookies = jar.make_cookies(response, request)
return cookies

View File

@ -1,58 +0,0 @@
import operator
from itertools import groupby
from collections import defaultdict
from pydispatch import dispatcher
from scrapy.core import signals
from scrapy.http import Response
from scrapy.utils.cookies import CookieJar
from scrapy.core.exceptions import HttpException
from scrapy import log
class CookiesMiddleware(object):
"""This middleware enables working with sites that need cookies"""
def __init__(self):
self.jars = defaultdict(CookieJar)
dispatcher.connect(self.domain_closed, signals.domain_closed)
def process_request(self, request, spider):
if request.meta.get('dont_merge_cookies', False):
return
jar = self.jars[spider.domain_name]
cookies = _get_cookies(jar, request)
for cookie in cookies:
jar.set_cookie_if_ok(cookie, request)
# set Cookie header
jar.add_cookie_header(request)
def process_response(self, request, response, spider):
if request.meta.get('dont_merge_cookies', False):
return
# extract cookies from Set-Cookie and drop invalid/expired cookies
jar = self.jars[spider.domain_name]
jar.extract_cookies(response, request)
# TODO: set current cookies in jar to response.cookies?
return response
# cookies should be set on non-200 responses too
def process_exception(self, request, exception, spider):
if isinstance(exception, HttpException):
self.process_response(request, exception.response, spider)
def domain_closed(self, domain):
self.jars.pop(domain, None)
def _get_cookies(jar, request):
headers = {'Set-Cookie': ['%s=%s;' % (k, v) for k, v in request.cookies.iteritems()]}
response = Response(request.url, headers=headers)
cookies = jar.make_cookies(response, request)
return cookies

View File

@ -55,7 +55,6 @@ def create_factory(request, spider):
postdata=request.body or None, # see http://dev.scrapy.org/ticket/60
headers=request.headers,
agent=agent,
cookies=request.cookies,
timeout=getattr(spider, "download_timeout", None) or default_timeout,
followRedirect=False)

View File

@ -1,8 +1,13 @@
from urlparse import urlunparse
from twisted.web.client import HTTPClientFactory
from twisted.web.client import HTTPClientFactory, HTTPPageGetter
from twisted.web import http
from twisted.python import failure
from twisted.web import error
from twisted.internet import defer
from scrapy.http import Url
from scrapy.http import Url, Headers
from scrapy.utils.misc import arg_to_iter
def _parse(url, defaultPort=None):
url = url.strip()
@ -27,6 +32,93 @@ def _parse(url, defaultPort=None):
return scheme, host, port, path
class ScrapyHTTPPageGetter(HTTPPageGetter):
quietLoss = 0
failed = 0
_specialHeaders = set(('host', 'user-agent', 'content-length'))
def connectionMade(self):
headers = self.factory.headers
method = getattr(self.factory, 'method', 'GET')
self.sendCommand(method, self.factory.path)
self.sendHeader('Host', headers.get("host", self.factory.host))
self.sendHeader('User-Agent', headers.get('User-Agent', self.factory.agent))
data = getattr(self.factory, 'postdata', None)
if data is not None:
self.sendHeader("Content-Length", str(len(data)))
for key, value in self.factory.headers.items():
if key.lower() not in self._specialHeaders:
self.sendHeader(key, value)
self.endHeaders()
self.headers = Headers()
if data is not None:
self.transport.write(data)
def sendHeader(self, name, value):
for v in arg_to_iter(value):
self.transport.write('%s: %s\r\n' % (name, v))
def handleHeader(self, key, value):
self.headers.appendlist(key, value)
def handleStatus(self, version, status, message):
self.version, self.status, self.message = version, status, message
self.factory.gotStatus(version, status, message)
def handleEndHeaders(self):
self.factory.gotHeaders(self.headers)
m = getattr(self, 'handleStatus_'+self.status, self.handleStatusDefault)
m()
def handleStatus_200(self):
pass
handleStatus_201 = lambda self: self.handleStatus_200()
handleStatus_202 = lambda self: self.handleStatus_200()
def handleStatusDefault(self):
self.failed = 1
def connectionLost(self, reason):
if not self.quietLoss:
http.HTTPClient.connectionLost(self, reason)
self.factory.noPage(reason)
def handleResponse(self, response):
if self.quietLoss:
return
if self.failed:
self.factory.noPage(failure.Failure(error.Error(self.status, self.message, response)))
if self.factory.method.upper() == 'HEAD':
# Callback with empty string, since there is never a response
# body for HEAD requests.
self.factory.page('')
elif self.length != None and self.length != 0:
self.factory.noPage(failure.Failure(
PartialDownloadError(self.status, self.message, response)))
else:
self.factory.page(response)
# server might be stupid and not close connection. admittedly
# the fact we do only one request per connection is also
# stupid...
self.transport.loseConnection()
def timeout(self):
self.quietLoss = True
self.transport.loseConnection()
self.factory.noPage(defer.TimeoutError("Getting %s took longer than %s seconds." % (self.factory.url, self.factory.timeout)))
class ScrapyHTTPClientFactory(HTTPClientFactory):
"""Scrapy implementation of the HTTPClientFactory overwriting the
serUrl method to make use of our Url object that cache the parse
@ -34,6 +126,8 @@ class ScrapyHTTPClientFactory(HTTPClientFactory):
cookies.
"""
protocol = ScrapyHTTPPageGetter
def setURL(self, url):
self.url = url
scheme, host, port, path = _parse(url)
@ -55,18 +149,4 @@ class ScrapyHTTPClientFactory(HTTPClientFactory):
values that could be managed correctly by twisted.
"""
self.response_headers = headers
if 'set-cookie' in headers:
goodcookies = []
for cookie in headers['set-cookie']:
cookparts = cookie.split(';')
cook = cookparts[0].lstrip()
t = cook.split('=', 1)
if len(t) == 2: #Good cookie
goodcookies.append(cookie)
k, v = t
self.cookies[k.lstrip()] = v.lstrip()
if goodcookies:
self.response_headers['set-cookie'] = goodcookies
else:
del self.response_headers['set-cookie']

View File

@ -48,7 +48,7 @@ class Headers(CaselessDict):
self[key] = list_
def setlistdefault(self, key, default_list=()):
self.setdefault(key, default_list)
return self.setdefault(key, default_list)
def appendlist(self, key, value):
lst = self.getlist(key)
@ -59,19 +59,16 @@ class Headers(CaselessDict):
return list(self.iteritems())
def iteritems(self):
return ((k, self[k]) for k in self.keys())
return ((k, self.getlist(k)) for k in self.keys())
def values(self):
return [self[k] for k in self.keys()]
def lists(self):
return super(Headers, self).items()
def to_string(self):
return headers_dict_to_raw(self)
def __copy__(self):
return self.__class__(self.lists())
return self.__class__(self)
copy = __copy__

View File

@ -60,13 +60,11 @@ class HeadersTest(unittest.TestCase):
h = Headers(idict)
self.assertEqual(dict(h), {'Content-Type': ['text/html'], 'X-Forwarded-For': ['ip1', 'ip2']})
self.assertEqual(h.keys(), ['X-Forwarded-For', 'Content-Type'])
self.assertEqual(h.items(), [('X-Forwarded-For', 'ip2'), ('Content-Type', 'text/html')])
self.assertEqual(h.items(), [('X-Forwarded-For', ['ip1', 'ip2']), ('Content-Type', ['text/html'])])
self.assertEqual(list(h.iteritems()),
[('X-Forwarded-For', 'ip2'), ('Content-Type', 'text/html')])
[('X-Forwarded-For', ['ip1', 'ip2']), ('Content-Type', ['text/html'])])
self.assertEqual(h.values(), ['ip2', 'text/html'])
self.assertEqual(h.lists(),
[('X-Forwarded-For', ['ip1', 'ip2']), ('Content-Type', ['text/html'])])
def test_update(self):
h = Headers()

View File

@ -52,7 +52,8 @@ class RequestTest(unittest.TestCase):
h[u'newkey'] = u'newval'
for k, v in h.iteritems():
self.assert_(isinstance(k, str))
self.assert_(isinstance(v, str))
for s in v:
self.assert_(isinstance(s, str))
def test_eq(self):
url = 'http://www.scrapy.org'

View File

@ -1081,6 +1081,23 @@ class LWPCookieTests(TestCase):
# we didn't have session cookies in the first place
counter["session_before"] == 0))
def testMalformedCookieHeaderParsing(self):
headers = Headers({'Set-Cookie': [
'CUSTOMER=WILE_E_COYOTE; path=/; expires=Wednesday, 09-Nov-2100 23:12:40 GMT',
'PART_NUMBER=ROCKET_LAUNCHER_0001; path=/',
'SHIPPING=FEDEX; path=/foo',
'COUNTRY=UY; path=/foo',
'GOOD_CUSTOMER;',
'NO_A_BOT;']})
res = Response('http://www.perlmeister.com/foo', headers=headers)
req = Request('http://www.perlmeister.com/foo')
c = CookieJar()
c.extract_cookies(res, req)
c.add_cookie_header(req)
self.assertEquals(req.headers.get('Cookie'),
'COUNTRY=UY; SHIPPING=FEDEX; CUSTOMER=WILE_E_COYOTE; '
'PART_NUMBER=ROCKET_LAUNCHER_0001; NO_A_BOT; GOOD_CUSTOMER')
class WrapperRequestResponse(TestCase):
@ -1130,3 +1147,4 @@ class WrapperRequestResponse(TestCase):
req = self.FakeRequest("http://www.acme.com:2345/resource.html", headers={"Host": "www.acme.com:5432"})
self.assertEquals(request_host(req), "www.acme.com")

View File

@ -1,4 +1,5 @@
"""
from twisted.internet import defer
Tests borrowed from the twisted.web.client tests.
"""
from urlparse import urlparse
@ -52,61 +53,3 @@ class ParseUrlTestCase(unittest.TestCase):
self.assertTrue(isinstance(path, str))
class FakeTransport:
disconnecting = False
def __init__(self):
self.data = []
def write(self, stuff):
self.data.append(stuff)
class CookieTestCase(unittest.TestCase):
def _listen(self, site):
return reactor.listenTCP(0, site, interface="127.0.0.1")
def setUp(self):
root = static.Data('El toro!', 'text/plain')
site = server.Site(root, timeout=None)
self.port = self._listen(site)
self.portno = self.port.getHost().port
def tearDown(self):
return self.port.stopListening()
def testMalformedCookieHeaderParsing(self):
d = defer.Deferred()
factory = ScrapyHTTPClientFactory('http://foo.example.com/')
proto = factory.buildProtocol('127.42.42.42')
proto.transport = FakeTransport()
proto.connectionMade()
for line in [
'200 Ok',
'Squash: yes',
'Hands: stolen',
'Set-Cookie: CUSTOMER=WILE_E_COYOTE; path=/; expires=Wednesday, 09-Nov-99 23:12:40 GMT',
'Set-Cookie: PART_NUMBER=ROCKET_LAUNCHER_0001; path=/',
'Set-Cookie: SHIPPING=FEDEX; path=/foo',
'Set-Cookie: COUNTRY=UY; path=/foo',
'Set-Cookie: GOOD_CUSTOMER;',
'Set-Cookie: NO_A_BOT;',
'',
'body',
'more body',
]:
proto.dataReceived(line + '\r\n')
self.assertEquals(proto.transport.data,
['GET / HTTP/1.0\r\n',
'Host: foo.example.com\r\n',
'User-Agent: Twisted PageGetter\r\n',
'\r\n'])
self.assertEquals(factory.cookies,
{
'CUSTOMER': 'WILE_E_COYOTE',
'PART_NUMBER': 'ROCKET_LAUNCHER_0001',
'SHIPPING': 'FEDEX',
'COUNTRY': 'UY',
})

View File

@ -95,7 +95,7 @@ class WrappedRequest(object):
return self.request.headers.items()
def add_unredirected_header(self, name, value):
self.request.headers[name] = value
self.request.headers.appendlist(name, value)
#print 'add_unredirected_header', self.request.headers

View File

@ -52,7 +52,8 @@ def request_fingerprint(request, include_headers=()):
for hdr in include_headers:
if hdr in request.headers:
fp.update(hdr)
fp.update(request.headers.get(hdr, ''))
for v in request.headers.getlist(hdr):
fp.update(v)
fphash = fp.hexdigest()
request.cache[cachekey] = fphash
return fphash