From 8dbbbd13950dcb21dda759b073c64ffdca85c2d6 Mon Sep 17 00:00:00 2001 From: Stas Glubokiy Date: Mon, 3 Sep 2018 20:07:37 +0300 Subject: [PATCH] Use request_cls attribute in contract definition --- docs/topics/contracts.rst | 5 +++-- scrapy/contracts/__init__.py | 15 +++++++++------ tests/test_contracts.py | 2 +- 3 files changed, 13 insertions(+), 9 deletions(-) diff --git a/docs/topics/contracts.rst b/docs/topics/contracts.rst index ada6fd227..70f20d4ed 100644 --- a/docs/topics/contracts.rst +++ b/docs/topics/contracts.rst @@ -86,8 +86,9 @@ override three methods: .. method:: Contract.adjust_request_args(args) This receives a ``dict`` as an argument containing default arguments - for request object. :class:`~scrapy.http.Request` is used - if ``request_cls`` is not set on ``args``. + for request object. :class:`~scrapy.http.Request` is used by default, + but this can be changed with the ``request_cls`` attribute. + If multiple contracts in chain have this attribute defined, the last one is used. Must return the same or a modified version of it. diff --git a/scrapy/contracts/__init__.py b/scrapy/contracts/__init__.py index 801c18e73..851a26a8e 100644 --- a/scrapy/contracts/__init__.py +++ b/scrapy/contracts/__init__.py @@ -4,7 +4,6 @@ from functools import wraps from inspect import getmembers from unittest import TestCase -from scrapy import FormRequest from scrapy.http import Request from scrapy.utils.spider import iterate_spider_output from scrapy.utils.python import get_spec @@ -50,14 +49,17 @@ class ContractsManager(object): def from_method(self, method, results): contracts = self.extract_contracts(method) if contracts: - # prepare request arguments - kwargs = {'callback': method} + request_cls = Request + for contract in contracts: + if contract.request_cls is not None: + request_cls = contract.request_cls + + # calculate request args + args, kwargs = get_spec(request_cls.__init__) + kwargs['callback'] = method for contract in contracts: kwargs = contract.adjust_request_args(kwargs) - request_cls = kwargs.pop('request_cls', Request) - - args, _ = get_spec(request_cls.__init__) args.remove('self') # check if all positional arguments are defined in kwargs @@ -98,6 +100,7 @@ class ContractsManager(object): class Contract(object): """ Abstract class for contracts """ + request_cls = None def __init__(self, method, *args): self.testcase_pre = _create_testcase(method, '@%s pre-hook' % self.name) diff --git a/tests/test_contracts.py b/tests/test_contracts.py index c35b068a4..fc5c94771 100644 --- a/tests/test_contracts.py +++ b/tests/test_contracts.py @@ -27,9 +27,9 @@ class ResponseMock(object): class CustomFormContract(Contract): name = 'custom_form' + request_cls = FormRequest def adjust_request_args(self, args): - args['request_cls'] = FormRequest args['formdata'] = {'name': 'scrapy'} return args