From debab6cff6fb147e7ce8f70c2e493c9568193397 Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?Adri=C3=A1n=20Chaves?= Date: Wed, 19 Mar 2025 20:16:49 +0100 Subject: [PATCH] Support defining Spider.start() as a sync generator --- scrapy/core/spidermw.py | 32 ++++++++++++++++++++++---------- tests/test_spider_start.py | 10 ++++++++++ 2 files changed, 32 insertions(+), 10 deletions(-) diff --git a/scrapy/core/spidermw.py b/scrapy/core/spidermw.py index 9a292cae9..95ce47394 100644 --- a/scrapy/core/spidermw.py +++ b/scrapy/core/spidermw.py @@ -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( diff --git a/tests/test_spider_start.py b/tests/test_spider_start.py index 34f1d5e66..5daa7d209 100644 --- a/tests/test_spider_start.py +++ b/tests/test_spider_start.py @@ -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):