diff --git a/src/crawlee/crawlers/_basic/_basic_crawler.py b/src/crawlee/crawlers/_basic/_basic_crawler.py index e6944ee2b0..0b9219fae0 100644 --- a/src/crawlee/crawlers/_basic/_basic_crawler.py +++ b/src/crawlee/crawlers/_basic/_basic_crawler.py @@ -1180,7 +1180,7 @@ async def _handle_request_retries( await self._mark_request_as_handled(request) return - await request_manager.reclaim_request(request) + await request_manager.reclaim_request(request, forefront=request.forefront) else: request.state = RequestState.ERROR await self._mark_request_as_handled(request) @@ -1474,22 +1474,37 @@ async def __run_task_function(self) -> None: if not session: raise RuntimeError('SessionError raised in a crawling context without a session') from session_error + request.state = RequestState.ERROR_HANDLER + + new_request = None if self._error_handler: - await self._error_handler(context, session_error) + try: + new_request = await self._error_handler(context, session_error) + except Exception as e: + raise UserDefinedErrorHandlerError('Exception thrown in user-defined request error handler') from e if self._should_retry_request(context, session_error): + # Replacement requests are only honored while rotations remain, so exhausted sessions + # still go through failed_request_handler instead of being silently replaced. + if new_request is not None and new_request != request: + await self._statistics.error_tracker_retry.add(error=session_error, context=context) + await request_manager.add_request(new_request) + await self._mark_request_as_handled(request) + session.retire() + return + exc_only = ''.join(traceback.format_exception_only(session_error)).strip() self._logger.warning('Encountered "%s", rotating session and retrying...', exc_only) - if session: - session.retire() + session.retire() # Increment session rotation count. request.session_rotation_count = (request.session_rotation_count or 0) + 1 - await request_manager.reclaim_request(request) + await request_manager.reclaim_request(request, forefront=request.forefront) await self._statistics.error_tracker_retry.add(error=session_error, context=context) else: + request.state = RequestState.ERROR await self._mark_request_as_handled(request) await self._handle_failed_request(context, session_error) diff --git a/tests/unit/crawlers/_basic/test_basic_crawler.py b/tests/unit/crawlers/_basic/test_basic_crawler.py index 763969b3f9..2643ebc897 100644 --- a/tests/unit/crawlers/_basic/test_basic_crawler.py +++ b/tests/unit/crawlers/_basic/test_basic_crawler.py @@ -274,6 +274,97 @@ async def error_handler(context: BasicCrawlingContext, error: Exception) -> None assert error_handler_mock.call_count == 1 +async def test_session_error_handler_can_replace_request() -> None: + """`error_handler` return value must be honored for SessionError while rotations remain.""" + queue = await RequestQueue.open() + crawler = BasicCrawler(request_manager=queue, max_session_rotations=3) + + request = Request.from_url('https://a.placeholder.com') + + @crawler.router.default_handler + async def handler(context: BasicCrawlingContext) -> None: + if '|recovered' in context.request.unique_key: + return + raise SessionError('blocked') + + @crawler.error_handler + async def error_handler(context: BasicCrawlingContext, error: Exception) -> Request | None: + return Request.from_url( + context.request.url, + unique_key=f'{context.request.unique_key}|recovered', + ) + + await crawler.run([request]) + + original_request = await queue.get_request(request.unique_key) + recovered_request = await queue.get_request(f'{request.unique_key}|recovered') + + assert original_request is not None + assert original_request.was_already_handled + assert recovered_request is not None + assert recovered_request.state == RequestState.DONE + assert recovered_request.was_already_handled + + await queue.drop() + + +async def test_session_error_handler_replacement_ignored_when_rotations_exhausted() -> None: + """When rotations are exhausted, replacement requests must not skip failed_request_handler.""" + queue = await RequestQueue.open() + failed_handler_mock = AsyncMock() + crawler = BasicCrawler(request_manager=queue, max_session_rotations=1) + + request = Request.from_url('https://a.placeholder.com') + + @crawler.router.default_handler + async def handler(context: BasicCrawlingContext) -> None: + raise SessionError('blocked') + + @crawler.error_handler + async def error_handler(context: BasicCrawlingContext, error: Exception) -> Request | None: + return Request.from_url( + context.request.url, + unique_key=f'{context.request.unique_key}|should-not-run', + ) + + @crawler.failed_request_handler + async def failed_request_handler(context: BasicCrawlingContext, error: Exception) -> None: + await failed_handler_mock(context, error) + + await crawler.run([request]) + + failed_handler_mock.assert_awaited_once() + assert await queue.get_request(f'{request.unique_key}|should-not-run') is None + original_request = await queue.get_request(request.unique_key) + assert original_request is not None + assert original_request.state == RequestState.ERROR + assert original_request.was_already_handled + + await queue.drop() + + +async def test_reclaim_uses_request_forefront_flag() -> None: + """Retries must reclaim with `request.forefront` so tiered-proxy priority retries stay at the front.""" + queue = await RequestQueue.open() + crawler = BasicCrawler(request_manager=queue, max_request_retries=1) + + @crawler.router.default_handler + async def handler(context: BasicCrawlingContext) -> None: + context.request.forefront = True + raise RuntimeError('Arbitrary crash for testing purposes') + + with patch.object(queue, 'reclaim_request', wraps=queue.reclaim_request) as reclaim_mock: + await crawler.run(['https://a.placeholder.com']) + + reclaim_mock.assert_awaited_once() + + (reclaimed_request,), reclaim_kwargs = reclaim_mock.await_args_list[0] + assert reclaimed_request.url == 'https://a.placeholder.com' + assert reclaim_kwargs == {'forefront': True} + + await queue.drop() + + async def test_handles_error_in_error_handler() -> None: crawler = BasicCrawler(max_request_retries=3)