Add typing for middleware and coroutine related code.

This commit is contained in:
Andrey Rakhmatullin 2021-04-03 17:40:45 +05:00
parent 8c5a3a5189
commit a9e96f9907
6 changed files with 56 additions and 35 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)):
@ -45,7 +49,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):
@ -64,7 +68,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,25 +3,32 @@ Spider Middleware manager
See documentation in docs/topics/spider-middleware.rst
"""
from collections.abc import Iterable, AsyncIterable
from itertools import islice
from typing import Callable, Union, Any
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__')
def _isiterable(o):
return isinstance(o, Iterable)
def _fname(f):
return f"{f.__self__.__class__.__name__}.{f.__func__.__name__}"
ScrapeFunc = Callable[[Union[Response, Failure], Request, Spider], Any]
class SpiderMiddlewareManager(MiddlewareManager):
component_name = 'spider middleware'
@ -41,7 +48,7 @@ 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):
for method in self.methods['process_spider_input']:
try:
result = method(response=response, spider=spider)
@ -55,7 +62,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):
try:
for r in iterable:
yield r
@ -66,7 +74,7 @@ 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=0):
exception = _failure.value
# don't handle _InvalidOutput exception
if isinstance(exception, _InvalidOutput):
@ -88,7 +96,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=0):
# 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()
@ -114,21 +123,21 @@ 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):
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):
def process_callback_output(result: Iterable):
return self._process_callback_output(response, spider, result)
def process_spider_exception(_failure):
def process_spider_exception(_failure: Failure):
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):
return self._process_chain('process_start_requests', start_requests, spider)

View File

@ -1,6 +1,7 @@
from collections import defaultdict, deque
import logging
import pprint
from collections import defaultdict, deque
from typing import Callable
from scrapy.exceptions import NotConfigured
from scrapy.utils.misc import create_instance, load_object
@ -16,7 +17,7 @@ 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)
@ -58,13 +59,13 @@ class MiddlewareManager:
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):
return process_parallel(self.methods[methodname], obj, *args)
def _process_chain(self, methodname, obj, *args):
def _process_chain(self, methodname: str, obj, *args):
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):
return process_chain_both(self.methods[cb_methodname],
self.methods[eb_methodname], obj, *args)

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,19 @@ Helper functions for dealing with Twisted deferreds
"""
import asyncio
import inspect
from collections.abc import Coroutine
from functools import wraps
from typing import Callable, Iterable, Any
from twisted.internet import defer, task
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):
"""Same as twisted.internet.defer.fail but delay calling errback until
next reactor loop
@ -47,7 +50,7 @@ def defer_result(result):
return defer_succeed(result)
def mustbe_deferred(f, *args, **kw):
def mustbe_deferred(f: Callable, *args, **kw):
"""Same as twisted.internet.defer.maybeDeferred, but delay calling
callback/errback to next reactor loop
"""
@ -64,7 +67,7 @@ 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):
"""Execute a callable over the objects in the given iterable, in parallel,
using no more than ``count`` concurrent calls.
@ -75,7 +78,7 @@ def parallel(iterable, count, callable, *args, **named):
return defer.DeferredList([coop.coiterate(work) for _ in range(count)])
def process_chain(callbacks, input, *a, **kw):
def process_chain(callbacks: Iterable[Callable], input, *a, **kw):
"""Return a Deferred built by chaining the given callbacks"""
d = defer.Deferred()
for x in callbacks:
@ -84,7 +87,7 @@ def process_chain(callbacks, input, *a, **kw):
return d
def process_chain_both(callbacks, errbacks, input, *a, **kw):
def process_chain_both(callbacks: Iterable[Callable], errbacks: Iterable[Callable], input, *a, **kw):
"""Return a Deferred built by chaining the given callbacks and errbacks"""
d = defer.Deferred()
for cb, eb in zip(callbacks, errbacks):
@ -100,7 +103,7 @@ 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):
"""Return a Deferred with the output of all successful calls to the given
callbacks
"""
@ -110,7 +113,7 @@ def process_parallel(callbacks, input, *a, **kw):
return d
def iter_errback(iterable, errback, *a, **kw):
def iter_errback(iterable: Iterable, errback: Callable, *a, **kw):
"""Wraps an iterable calling an errback if an error is caught while
iterating it.
"""
@ -124,7 +127,7 @@ 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):
return o
@ -139,7 +142,7 @@ def deferred_from_coro(o):
return o
def deferred_f_from_coro_f(coro_f):
def deferred_f_from_coro_f(coro_f: Callable[..., Coroutine]):
""" 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,7 +154,7 @@ def deferred_f_from_coro_f(coro_f):
return f
def maybeDeferred_coro(f, *args, **kw):
def maybeDeferred_coro(f: Callable, *args, **kw):
""" Copy of defer.maybeDeferred that also converts coroutines to Deferreds. """
try:
result = f(*args, **kw)

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):