mirror of https://github.com/scrapy/scrapy.git
Support defining Spider.start() as a sync generator
This commit is contained in:
parent
ea3e6d2e42
commit
debab6cff6
|
|
@ -8,7 +8,12 @@ from __future__ import annotations
|
|||
|
||||
import logging
|
||||
from collections.abc import AsyncIterable, Callable, Iterable
|
||||
from inspect import isasyncgenfunction, iscoroutine, iscoroutinefunction
|
||||
from functools import wraps
|
||||
from inspect import (
|
||||
isasyncgenfunction,
|
||||
iscoroutine,
|
||||
isgeneratorfunction,
|
||||
)
|
||||
from itertools import islice
|
||||
from typing import TYPE_CHECKING, Any, TypeVar, Union, cast
|
||||
from warnings import warn
|
||||
|
|
@ -49,6 +54,21 @@ def _isiterable(o: Any) -> bool:
|
|||
return isinstance(o, (Iterable, AsyncIterable))
|
||||
|
||||
|
||||
def _sync_generator_to_async(f: Callable) -> Callable:
|
||||
@wraps(f)
|
||||
async def wrapper(*args, **kwargs):
|
||||
for item in f(*args, **kwargs):
|
||||
yield item
|
||||
|
||||
return wrapper
|
||||
|
||||
|
||||
def _maybe_sync_generator_to_async(f: Callable) -> Callable:
|
||||
if isgeneratorfunction(f):
|
||||
return _sync_generator_to_async(f)
|
||||
return f
|
||||
|
||||
|
||||
class SpiderMiddlewareManager(MiddlewareManager):
|
||||
component_name = "spider middleware"
|
||||
|
||||
|
|
@ -380,7 +400,7 @@ class SpiderMiddlewareManager(MiddlewareManager):
|
|||
)
|
||||
seeds = as_async_generator(sync_seeds)
|
||||
else:
|
||||
seeds = yield self._iter_seeds(spider)
|
||||
seeds = yield _maybe_sync_generator_to_async(spider.start)()
|
||||
seeds = yield self._process_chain("process_start", seeds)
|
||||
return seeds
|
||||
|
||||
|
|
@ -454,14 +474,6 @@ class SpiderMiddlewareManager(MiddlewareManager):
|
|||
f"https://docs.scrapy.org/en/VERSION/news.html"
|
||||
)
|
||||
|
||||
@staticmethod
|
||||
def _iter_seeds(spider: Spider):
|
||||
fn = spider.start
|
||||
if isasyncgenfunction(fn):
|
||||
return fn().__aiter__()
|
||||
assert iscoroutinefunction(fn)
|
||||
return deferred_from_coro(fn())
|
||||
|
||||
# This method is only needed until _async compatibility methods are removed.
|
||||
@staticmethod
|
||||
def _get_async_method_pair(
|
||||
|
|
|
|||
|
|
@ -75,6 +75,16 @@ class MainTestCase(TestCase):
|
|||
|
||||
await self._test_spider(TestSpider, [ITEM_A])
|
||||
|
||||
@deferred_f_from_coro_f
|
||||
async def test_start_sync(self):
|
||||
class TestSpider(Spider):
|
||||
name = "test"
|
||||
|
||||
def start(self):
|
||||
yield ITEM_A
|
||||
|
||||
await self._test_spider(TestSpider, [ITEM_A])
|
||||
|
||||
@deferred_f_from_coro_f
|
||||
async def test_deprecated(self):
|
||||
class TestSpider(Spider):
|
||||
|
|
|
|||
Loading…
Reference in New Issue