Merge pull request #5077 from wRAR/deferred-typing

Add typing for middleware and coroutine related code.
This commit is contained in:
Andrey Rahmatullin 2021-04-13 23:29:49 +05:00 committed by GitHub
commit 9bf9ab7291
No known key found for this signature in database
GPG Key ID: 4AEE18F83AFDEB23
7 changed files with 94 additions and 62 deletions

View File

@ -3,8 +3,12 @@ Downloader Middleware manager
See documentation in docs/topics/downloader-middleware.rst
"""
from twisted.internet import defer
from typing import Callable, Union
from twisted.internet import defer
from twisted.python.failure import Failure
from scrapy import Spider
from scrapy.exceptions import _InvalidOutput
from scrapy.http import Request, Response
from scrapy.middleware import MiddlewareManager
@ -29,9 +33,9 @@ class DownloaderMiddlewareManager(MiddlewareManager):
if hasattr(mw, 'process_exception'):
self.methods['process_exception'].appendleft(mw.process_exception)
def download(self, download_func, request, spider):
def download(self, download_func: Callable, request: Request, spider: Spider):
@defer.inlineCallbacks
def process_request(request):
def process_request(request: Request):
for method in self.methods['process_request']:
response = yield deferred_from_coro(method(request=request, spider=spider))
if response is not None and not isinstance(response, (Response, Request)):
@ -44,7 +48,7 @@ class DownloaderMiddlewareManager(MiddlewareManager):
return (yield download_func(request=request, spider=spider))
@defer.inlineCallbacks
def process_response(response):
def process_response(response: Union[Response, Request]):
if response is None:
raise TypeError("Received None in process_response")
elif isinstance(response, Request):
@ -62,7 +66,7 @@ class DownloaderMiddlewareManager(MiddlewareManager):
return response
@defer.inlineCallbacks
def process_exception(failure):
def process_exception(failure: Failure):
exception = failure.value
for method in self.methods['process_exception']:
response = yield deferred_from_coro(method(request=request, exception=exception, spider=spider))

View File

@ -3,12 +3,14 @@ extracts information from them"""
import logging
from collections import deque
from collections.abc import Iterable
from typing import Union
from itemadapter import is_item
from twisted.internet import defer
from twisted.python.failure import Failure
from scrapy import signals
from scrapy import signals, Spider
from scrapy.core.spidermw import SpiderMiddlewareManager
from scrapy.exceptions import CloseSpider, DropItem, IgnoreRequest
from scrapy.http import Request, Response
@ -118,7 +120,7 @@ class Scraper:
response, request, deferred = self.slot.next_response_request_deferred()
self._scrape(response, request, spider).chainDeferred(deferred)
def _scrape(self, result, request, spider):
def _scrape(self, result: Union[Response, Failure], request: Request, spider: Spider):
"""
Handle the downloaded response or failure through the spider callback/errback
"""
@ -129,7 +131,7 @@ class Scraper:
dfd.addCallback(self.handle_spider_output, request, result, spider)
return dfd
def _scrape2(self, result, request, spider):
def _scrape2(self, result: Union[Response, Failure], request: Request, spider: Spider):
"""
Handle the different cases of request's result been a Response or a Failure
"""
@ -139,7 +141,7 @@ class Scraper:
dfd = self.call_spider(result, request, spider)
return dfd.addErrback(self._log_download_errors, result, request, spider)
def call_spider(self, result, request, spider):
def call_spider(self, result: Union[Response, Failure], request: Request, spider: Spider):
if isinstance(result, Response):
if getattr(result, "request", None) is None:
result.request = request
@ -154,7 +156,7 @@ class Scraper:
dfd.addErrback(request.errback)
return dfd.addCallback(iterate_spider_output)
def handle_spider_error(self, _failure, request, response, spider):
def handle_spider_error(self, _failure: Failure, request: Request, response: Response, spider: Spider):
exc = _failure.value
if isinstance(exc, CloseSpider):
self.crawler.engine.close_spider(spider, exc.reason or 'cancelled')
@ -175,7 +177,7 @@ class Scraper:
spider=spider
)
def handle_spider_output(self, result, request, response, spider):
def handle_spider_output(self, result: Iterable, request: Request, response: Response, spider: Spider):
if not result:
return defer_succeed(None)
it = iter_errback(result, self.handle_spider_error, request, response, spider)

View File

@ -4,18 +4,25 @@ Spider Middleware manager
See documentation in docs/topics/spider-middleware.rst
"""
from itertools import islice
from typing import Any, Callable, Generator, Iterable, Union
from twisted.internet.defer import Deferred
from twisted.python.failure import Failure
from scrapy import Request, Spider
from scrapy.exceptions import _InvalidOutput
from scrapy.http import Response
from scrapy.middleware import MiddlewareManager
from scrapy.utils.conf import build_component_list
from scrapy.utils.defer import mustbe_deferred
from scrapy.utils.python import MutableChain
def _isiterable(possible_iterator):
return hasattr(possible_iterator, '__iter__')
ScrapeFunc = Callable[[Union[Response, Failure], Request, Spider], Any]
def _isiterable(o) -> bool:
return isinstance(o, Iterable)
class SpiderMiddlewareManager(MiddlewareManager):
@ -37,7 +44,8 @@ class SpiderMiddlewareManager(MiddlewareManager):
process_spider_exception = getattr(mw, 'process_spider_exception', None)
self.methods['process_spider_exception'].appendleft(process_spider_exception)
def _process_spider_input(self, scrape_func, response, request, spider):
def _process_spider_input(self, scrape_func: ScrapeFunc, response: Response, request: Request,
spider: Spider) -> Any:
for method in self.methods['process_spider_input']:
try:
result = method(response=response, spider=spider)
@ -51,7 +59,8 @@ class SpiderMiddlewareManager(MiddlewareManager):
return scrape_func(Failure(), request, spider)
return scrape_func(response, request, spider)
def _evaluate_iterable(self, response, spider, iterable, exception_processor_index, recover_to):
def _evaluate_iterable(self, response: Response, spider: Spider, iterable: Iterable,
exception_processor_index: int, recover_to: MutableChain) -> Generator:
try:
for r in iterable:
yield r
@ -62,7 +71,8 @@ class SpiderMiddlewareManager(MiddlewareManager):
raise
recover_to.extend(exception_result)
def _process_spider_exception(self, response, spider, _failure, start_index=0):
def _process_spider_exception(self, response: Response, spider: Spider, _failure: Failure,
start_index: int = 0) -> Union[Failure, MutableChain]:
exception = _failure.value
# don't handle _InvalidOutput exception
if isinstance(exception, _InvalidOutput):
@ -84,7 +94,8 @@ class SpiderMiddlewareManager(MiddlewareManager):
raise _InvalidOutput(msg)
return _failure
def _process_spider_output(self, response, spider, result, start_index=0):
def _process_spider_output(self, response: Response, spider: Spider,
result: Iterable, start_index: int = 0) -> MutableChain:
# items in this iterable do not need to go through the process_spider_output
# chain, they went through it already from the process_spider_exception method
recovered = MutableChain()
@ -110,21 +121,22 @@ class SpiderMiddlewareManager(MiddlewareManager):
return MutableChain(result, recovered)
def _process_callback_output(self, response, spider, result):
def _process_callback_output(self, response: Response, spider: Spider, result: Iterable) -> MutableChain:
recovered = MutableChain()
result = self._evaluate_iterable(response, spider, result, 0, recovered)
return MutableChain(self._process_spider_output(response, spider, result), recovered)
def scrape_response(self, scrape_func, response, request, spider):
def process_callback_output(result):
def scrape_response(self, scrape_func: ScrapeFunc, response: Response, request: Request,
spider: Spider) -> Deferred:
def process_callback_output(result: Iterable) -> MutableChain:
return self._process_callback_output(response, spider, result)
def process_spider_exception(_failure):
def process_spider_exception(_failure: Failure) -> Union[Failure, MutableChain]:
return self._process_spider_exception(response, spider, _failure)
dfd = mustbe_deferred(self._process_spider_input, scrape_func, response, request, spider)
dfd.addCallbacks(callback=process_callback_output, errback=process_spider_exception)
return dfd
def process_start_requests(self, start_requests, spider):
def process_start_requests(self, start_requests, spider: Spider) -> Deferred:
return self._process_chain('process_start_requests', start_requests, spider)

View File

@ -1,8 +1,13 @@
from collections import defaultdict, deque
import logging
import pprint
from collections import defaultdict, deque
from typing import Callable, Deque, Dict
from twisted.internet.defer import Deferred
from scrapy import Spider
from scrapy.exceptions import NotConfigured
from scrapy.settings import Settings
from scrapy.utils.misc import create_instance, load_object
from scrapy.utils.defer import process_parallel, process_chain, process_chain_both
@ -16,16 +21,16 @@ class MiddlewareManager:
def __init__(self, *middlewares):
self.middlewares = middlewares
self.methods = defaultdict(deque)
self.methods: Dict[str, Deque[Callable]] = defaultdict(deque)
for mw in middlewares:
self._add_middleware(mw)
@classmethod
def _get_mwlist_from_settings(cls, settings):
def _get_mwlist_from_settings(cls, settings: Settings) -> list:
raise NotImplementedError
@classmethod
def from_settings(cls, settings, crawler=None):
def from_settings(cls, settings: Settings, crawler=None):
mwlist = cls._get_mwlist_from_settings(settings)
middlewares = []
enabled = []
@ -52,24 +57,24 @@ class MiddlewareManager:
def from_crawler(cls, crawler):
return cls.from_settings(crawler.settings, crawler)
def _add_middleware(self, mw):
def _add_middleware(self, mw) -> None:
if hasattr(mw, 'open_spider'):
self.methods['open_spider'].append(mw.open_spider)
if hasattr(mw, 'close_spider'):
self.methods['close_spider'].appendleft(mw.close_spider)
def _process_parallel(self, methodname, obj, *args):
def _process_parallel(self, methodname: str, obj, *args) -> Deferred:
return process_parallel(self.methods[methodname], obj, *args)
def _process_chain(self, methodname, obj, *args):
def _process_chain(self, methodname: str, obj, *args) -> Deferred:
return process_chain(self.methods[methodname], obj, *args)
def _process_chain_both(self, cb_methodname, eb_methodname, obj, *args):
def _process_chain_both(self, cb_methodname: str, eb_methodname: str, obj, *args) -> Deferred:
return process_chain_both(self.methods[cb_methodname],
self.methods[eb_methodname], obj, *args)
def open_spider(self, spider):
def open_spider(self, spider: Spider) -> Deferred:
return self._process_parallel('open_spider', spider)
def close_spider(self, spider):
def close_spider(self, spider: Spider) -> Deferred:
return self._process_parallel('close_spider', spider)

View File

@ -1,4 +1,7 @@
async def collect_asyncgen(result):
from collections.abc import AsyncIterable
async def collect_asyncgen(result: AsyncIterable):
results = []
async for x in result:
results.append(x)

View File

@ -3,16 +3,21 @@ Helper functions for dealing with Twisted deferreds
"""
import asyncio
import inspect
from collections.abc import Coroutine
from functools import wraps
from typing import Any, Callable, Generator, Iterable
from twisted.internet import defer, task
from twisted.internet import defer
from twisted.internet.defer import Deferred, DeferredList, ensureDeferred
from twisted.internet.task import Cooperator
from twisted.python import failure
from twisted.python.failure import Failure
from scrapy.exceptions import IgnoreRequest
from scrapy.utils.reactor import is_asyncio_reactor_installed
def defer_fail(_failure):
def defer_fail(_failure: Failure) -> Deferred:
"""Same as twisted.internet.defer.fail but delay calling errback until
next reactor loop
@ -20,12 +25,12 @@ def defer_fail(_failure):
before attending pending delayed calls, so do not set delay to zero.
"""
from twisted.internet import reactor
d = defer.Deferred()
d = Deferred()
reactor.callLater(0.1, d.errback, _failure)
return d
def defer_succeed(result):
def defer_succeed(result) -> Deferred:
"""Same as twisted.internet.defer.succeed but delay calling callback until
next reactor loop
@ -33,13 +38,13 @@ def defer_succeed(result):
before attending pending delayed calls, so do not set delay to zero.
"""
from twisted.internet import reactor
d = defer.Deferred()
d = Deferred()
reactor.callLater(0.1, d.callback, result)
return d
def defer_result(result):
if isinstance(result, defer.Deferred):
def defer_result(result) -> Deferred:
if isinstance(result, Deferred):
return result
elif isinstance(result, failure.Failure):
return defer_fail(result)
@ -47,7 +52,7 @@ def defer_result(result):
return defer_succeed(result)
def mustbe_deferred(f, *args, **kw):
def mustbe_deferred(f: Callable, *args, **kw) -> Deferred:
"""Same as twisted.internet.defer.maybeDeferred, but delay calling
callback/errback to next reactor loop
"""
@ -64,29 +69,29 @@ def mustbe_deferred(f, *args, **kw):
return defer_result(result)
def parallel(iterable, count, callable, *args, **named):
def parallel(iterable: Iterable, count: int, callable: Callable, *args, **named) -> DeferredList:
"""Execute a callable over the objects in the given iterable, in parallel,
using no more than ``count`` concurrent calls.
Taken from: https://jcalderone.livejournal.com/24285.html
"""
coop = task.Cooperator()
coop = Cooperator()
work = (callable(elem, *args, **named) for elem in iterable)
return defer.DeferredList([coop.coiterate(work) for _ in range(count)])
return DeferredList([coop.coiterate(work) for _ in range(count)])
def process_chain(callbacks, input, *a, **kw):
def process_chain(callbacks: Iterable[Callable], input, *a, **kw) -> Deferred:
"""Return a Deferred built by chaining the given callbacks"""
d = defer.Deferred()
d = Deferred()
for x in callbacks:
d.addCallback(x, *a, **kw)
d.callback(input)
return d
def process_chain_both(callbacks, errbacks, input, *a, **kw):
def process_chain_both(callbacks: Iterable[Callable], errbacks: Iterable[Callable], input, *a, **kw) -> Deferred:
"""Return a Deferred built by chaining the given callbacks and errbacks"""
d = defer.Deferred()
d = Deferred()
for cb, eb in zip(callbacks, errbacks):
d.addCallbacks(
callback=cb, errback=eb,
@ -100,17 +105,17 @@ def process_chain_both(callbacks, errbacks, input, *a, **kw):
return d
def process_parallel(callbacks, input, *a, **kw):
def process_parallel(callbacks: Iterable[Callable], input, *a, **kw) -> Deferred:
"""Return a Deferred with the output of all successful calls to the given
callbacks
"""
dfds = [defer.succeed(input).addCallback(x, *a, **kw) for x in callbacks]
d = defer.DeferredList(dfds, fireOnOneErrback=True, consumeErrors=True)
d = DeferredList(dfds, fireOnOneErrback=True, consumeErrors=True)
d.addCallbacks(lambda r: [x[1] for x in r], lambda f: f.value.subFailure)
return d
def iter_errback(iterable, errback, *a, **kw):
def iter_errback(iterable: Iterable, errback: Callable, *a, **kw) -> Generator:
"""Wraps an iterable calling an errback if an error is caught while
iterating it.
"""
@ -124,22 +129,22 @@ def iter_errback(iterable, errback, *a, **kw):
errback(failure.Failure(), *a, **kw)
def deferred_from_coro(o):
def deferred_from_coro(o) -> Any:
"""Converts a coroutine into a Deferred, or returns the object as is if it isn't a coroutine"""
if isinstance(o, defer.Deferred):
if isinstance(o, Deferred):
return o
if asyncio.isfuture(o) or inspect.isawaitable(o):
if not is_asyncio_reactor_installed():
# wrapping the coroutine directly into a Deferred, this doesn't work correctly with coroutines
# that use asyncio, e.g. "await asyncio.sleep(1)"
return defer.ensureDeferred(o)
return ensureDeferred(o)
else:
# wrapping the coroutine into a Future and then into a Deferred, this requires AsyncioSelectorReactor
return defer.Deferred.fromFuture(asyncio.ensure_future(o))
return Deferred.fromFuture(asyncio.ensure_future(o))
return o
def deferred_f_from_coro_f(coro_f):
def deferred_f_from_coro_f(coro_f: Callable[..., Coroutine]) -> Callable:
""" Converts a coroutine function into a function that returns a Deferred.
The coroutine function will be called at the time when the wrapper is called. Wrapper args will be passed to it.
@ -151,14 +156,14 @@ def deferred_f_from_coro_f(coro_f):
return f
def maybeDeferred_coro(f, *args, **kw):
def maybeDeferred_coro(f: Callable, *args, **kw) -> Deferred:
""" Copy of defer.maybeDeferred that also converts coroutines to Deferreds. """
try:
result = f(*args, **kw)
except: # noqa: E722
return defer.fail(failure.Failure(captureVars=defer.Deferred.debug))
return defer.fail(failure.Failure(captureVars=Deferred.debug))
if isinstance(result, defer.Deferred):
if isinstance(result, Deferred):
return result
elif asyncio.isfuture(result) or inspect.isawaitable(result):
return deferred_from_coro(result)

View File

@ -8,6 +8,7 @@ import re
import sys
import warnings
import weakref
from collections.abc import Iterable
from functools import partial, wraps
from itertools import chain
@ -335,15 +336,15 @@ else:
gc.collect()
class MutableChain:
class MutableChain(Iterable):
"""
Thin wrapper around itertools.chain, allowing to add iterables "in-place"
"""
def __init__(self, *args):
def __init__(self, *args: Iterable):
self.data = chain.from_iterable(args)
def extend(self, *iterables):
def extend(self, *iterables: Iterable):
self.data = chain(self.data, chain.from_iterable(iterables))
def __iter__(self):