Add typing also for return values, other small fixes.

This commit is contained in:
Andrey Rakhmatullin 2021-04-13 20:01:18 +05:00
parent 5b547f0808
commit 76fa2257ef
3 changed files with 44 additions and 37 deletions

View File

@ -3,10 +3,10 @@ Spider Middleware manager
See documentation in docs/topics/spider-middleware.rst
"""
from collections.abc import Iterable
from itertools import islice
from typing import Callable, Union, Any
from typing import Callable, Union, Any, Generator, Iterable
from twisted.internet.defer import Deferred
from twisted.python.failure import Failure
from scrapy import Request, Spider
@ -18,13 +18,13 @@ from scrapy.utils.defer import mustbe_deferred
from scrapy.utils.python import MutableChain
def _isiterable(o):
return isinstance(o, Iterable)
ScrapeFunc = Callable[[Union[Response, Failure], Request, Spider], Any]
def _isiterable(o) -> bool:
return isinstance(o, Iterable)
class SpiderMiddlewareManager(MiddlewareManager):
component_name = 'spider middleware'
@ -44,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: ScrapeFunc, response: Response, request: Request, spider: 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)
@ -59,7 +60,7 @@ class SpiderMiddlewareManager(MiddlewareManager):
return scrape_func(response, request, spider)
def _evaluate_iterable(self, response: Response, spider: Spider, iterable: Iterable,
exception_processor_index: int, recover_to: MutableChain):
exception_processor_index: int, recover_to: MutableChain) -> Generator:
try:
for r in iterable:
yield r
@ -70,7 +71,8 @@ class SpiderMiddlewareManager(MiddlewareManager):
raise
recover_to.extend(exception_result)
def _process_spider_exception(self, response: Response, spider: Spider, _failure: 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):
@ -93,7 +95,7 @@ class SpiderMiddlewareManager(MiddlewareManager):
return _failure
def _process_spider_output(self, response: Response, spider: Spider,
result: Iterable, start_index=0):
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()
@ -119,21 +121,22 @@ class SpiderMiddlewareManager(MiddlewareManager):
return MutableChain(result, recovered)
def _process_callback_output(self, response: Response, spider: Spider, result: Iterable):
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: ScrapeFunc, response: Response, request: Request, spider: Spider):
def process_callback_output(result: Iterable):
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: 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: Spider):
def process_start_requests(self, start_requests, spider: Spider) -> Deferred:
return self._process_chain('process_start_requests', start_requests, spider)

View File

@ -1,9 +1,13 @@
import logging
import pprint
from collections import defaultdict, deque
from typing import Callable
from typing import Callable, Dict, Deque
from twisted.internet import defer
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
@ -17,16 +21,16 @@ class MiddlewareManager:
def __init__(self, *middlewares):
self.middlewares = middlewares
self.methods: dict[str, deque[Callable]] = 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 = []
@ -53,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: str, obj, *args):
def _process_parallel(self, methodname: str, obj, *args) -> defer.Deferred:
return process_parallel(self.methods[methodname], obj, *args)
def _process_chain(self, methodname: str, obj, *args):
def _process_chain(self, methodname: str, obj, *args) -> defer.Deferred:
return process_chain(self.methods[methodname], obj, *args)
def _process_chain_both(self, cb_methodname: str, eb_methodname: str, obj, *args):
def _process_chain_both(self, cb_methodname: str, eb_methodname: str, obj, *args) -> defer.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) -> defer.Deferred:
return self._process_parallel('open_spider', spider)
def close_spider(self, spider):
def close_spider(self, spider: Spider) -> defer.Deferred:
return self._process_parallel('close_spider', spider)

View File

@ -5,7 +5,7 @@ import asyncio
import inspect
from collections.abc import Coroutine
from functools import wraps
from typing import Callable, Iterable, Any
from typing import Callable, Iterable, Any, Generator
from twisted.internet import defer, task
from twisted.python import failure
@ -15,7 +15,7 @@ from scrapy.exceptions import IgnoreRequest
from scrapy.utils.reactor import is_asyncio_reactor_installed
def defer_fail(_failure: Failure):
def defer_fail(_failure: Failure) -> defer.Deferred:
"""Same as twisted.internet.defer.fail but delay calling errback until
next reactor loop
@ -28,7 +28,7 @@ def defer_fail(_failure: Failure):
return d
def defer_succeed(result):
def defer_succeed(result) -> defer.Deferred:
"""Same as twisted.internet.defer.succeed but delay calling callback until
next reactor loop
@ -41,7 +41,7 @@ def defer_succeed(result):
return d
def defer_result(result):
def defer_result(result) -> defer.Deferred:
if isinstance(result, defer.Deferred):
return result
elif isinstance(result, failure.Failure):
@ -50,7 +50,7 @@ def defer_result(result):
return defer_succeed(result)
def mustbe_deferred(f: Callable, *args, **kw):
def mustbe_deferred(f: Callable, *args, **kw) -> defer.Deferred:
"""Same as twisted.internet.defer.maybeDeferred, but delay calling
callback/errback to next reactor loop
"""
@ -67,7 +67,7 @@ def mustbe_deferred(f: Callable, *args, **kw):
return defer_result(result)
def parallel(iterable: Iterable, count: int, callable: Callable, *args, **named):
def parallel(iterable: Iterable, count: int, callable: Callable, *args, **named) -> defer.DeferredList:
"""Execute a callable over the objects in the given iterable, in parallel,
using no more than ``count`` concurrent calls.
@ -78,7 +78,7 @@ def parallel(iterable: Iterable, count: int, callable: Callable, *args, **named)
return defer.DeferredList([coop.coiterate(work) for _ in range(count)])
def process_chain(callbacks: Iterable[Callable], input, *a, **kw):
def process_chain(callbacks: Iterable[Callable], input, *a, **kw) -> defer.Deferred:
"""Return a Deferred built by chaining the given callbacks"""
d = defer.Deferred()
for x in callbacks:
@ -87,7 +87,7 @@ def process_chain(callbacks: Iterable[Callable], input, *a, **kw):
return d
def process_chain_both(callbacks: Iterable[Callable], errbacks: Iterable[Callable], input, *a, **kw):
def process_chain_both(callbacks: Iterable[Callable], errbacks: Iterable[Callable], input, *a, **kw) -> defer.Deferred:
"""Return a Deferred built by chaining the given callbacks and errbacks"""
d = defer.Deferred()
for cb, eb in zip(callbacks, errbacks):
@ -103,7 +103,7 @@ def process_chain_both(callbacks: Iterable[Callable], errbacks: Iterable[Callabl
return d
def process_parallel(callbacks: Iterable[Callable], input, *a, **kw):
def process_parallel(callbacks: Iterable[Callable], input, *a, **kw) -> defer.Deferred:
"""Return a Deferred with the output of all successful calls to the given
callbacks
"""
@ -113,7 +113,7 @@ def process_parallel(callbacks: Iterable[Callable], input, *a, **kw):
return d
def iter_errback(iterable: Iterable, errback: Callable, *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.
"""
@ -142,7 +142,7 @@ def deferred_from_coro(o) -> Any:
return o
def deferred_f_from_coro_f(coro_f: Callable[..., Coroutine]):
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.
@ -154,7 +154,7 @@ def deferred_f_from_coro_f(coro_f: Callable[..., Coroutine]):
return f
def maybeDeferred_coro(f: Callable, *args, **kw):
def maybeDeferred_coro(f: Callable, *args, **kw) -> defer.Deferred:
""" Copy of defer.maybeDeferred that also converts coroutines to Deferreds. """
try:
result = f(*args, **kw)