Support defining Spider.start() as a sync generator

This commit is contained in:
Adrián Chaves 2025-03-19 20:16:49 +01:00
parent ea3e6d2e42
commit debab6cff6
2 changed files with 32 additions and 10 deletions

View File

@ -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(

View File

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