refactor: unify paged search task orchestration

This commit is contained in:
jxxghp
2026-08-24 10:15:02 +08:00
parent f55f5edeca
commit 73d9d0dbf5
3 changed files with 386 additions and 267 deletions
+250 -266
View File
@@ -5,8 +5,9 @@ import random
import re
import time
from concurrent.futures import FIRST_COMPLETED, ThreadPoolExecutor, as_completed, wait
from contextlib import aclosing
from datetime import datetime
from typing import AsyncIterator, Any, Dict, Iterable, Tuple
from typing import AsyncIterator, Any, Awaitable, Callable, Dict, Iterable, Tuple
from typing import List, Optional
from unicodedata import normalize
@@ -2420,6 +2421,68 @@ class SearchChain(ChainBase):
# 返回
return results
async def _iter_site_page_results(
self,
*,
indexer_sites: List[dict],
search_pages: List[int],
search_page: Callable[[dict, int], Awaitable[Optional[List[Any]]]],
should_continue: Callable[[dict, List[Any]], bool],
task_owner: str,
) -> AsyncIterator[Tuple[dict, int, List[Any], bool]]:
"""统一调度站点逐页请求,并在调用方退出时取消、等待全部请求。"""
total_num = len(indexer_sites) * len(search_pages)
semaphore = asyncio.Semaphore(
self.runtime_config.search_threadpool_size or max(1, total_num)
)
pending_tasks: dict[
asyncio.Task[List[Any]],
Tuple[dict, int, int],
] = {}
async def run_site_page(site: dict, page_number: int) -> List[Any]:
"""在共享并发预算内执行一页站点请求并规范化空结果。"""
async with semaphore:
return await search_page(site, page_number) or []
def submit_site_page(site: dict, page_index: int) -> None:
"""登记一页请求及其续页位置,供统一终态收口。"""
page_number = search_pages[page_index]
task = asyncio.create_task(
run_site_page(site, page_number),
name=task_owner,
)
pending_tasks[task] = (site, page_index, page_number)
for site in indexer_sites:
submit_site_page(site, 0)
try:
while pending_tasks:
if global_vars.is_system_stopped:
break
done_tasks, _ = await asyncio.wait(
pending_tasks,
return_when=asyncio.FIRST_COMPLETED,
)
for task in done_tasks:
site, page_index, page_number = pending_tasks.pop(task)
page_results = await task
continued = (
should_continue(site, page_results)
and page_index + 1 < len(search_pages)
)
if continued:
submit_site_page(site, page_index + 1)
yield site, page_number, page_results, continued
finally:
tasks = tuple(pending_tasks)
for task in tasks:
if not task.done():
task.cancel()
if tasks:
await asyncio.gather(*tasks, return_exceptions=True)
async def __async_search_all_sites(self, keyword: str,
mediainfo: Optional[MediaInfo] = None,
sites: List[int] = None,
@@ -2472,75 +2535,57 @@ class SearchChain(ChainBase):
text=f"开始搜索,共 {len(indexer_sites)} 个站点,{len(search_pages)} 页 ...")
# 结果集
results = list(plugin_results)
semaphore = asyncio.Semaphore(
self.runtime_config.search_threadpool_size or total_num
)
async def search_site_page(site: dict, search_page: int) -> List[TorrentInfo]:
"""
控制单次站点页请求的并发量,并返回该页的资源列表。
"""
async with semaphore:
if area == "imdbid":
# 搜索IMDBID
return await self.async_search_site_torrents(site=site,
keyword=mediainfo.imdb_id if mediainfo else None,
mtype=mediainfo.type if mediainfo else mtype,
page=search_page)
# 搜索标题
return await self.async_search_site_torrents(site=site,
keyword=keyword,
mtype=mediainfo.type if mediainfo else mtype,
page=search_page)
"""调用既有站点资源接口,具体并发由统一逐页编排器控制。"""
search_keyword = (
mediainfo.imdb_id
if area == "imdbid" and mediainfo
else keyword
)
return await self.async_search_site_torrents(
site=site,
keyword=search_keyword,
mtype=mediainfo.type if mediainfo else mtype,
page=search_page,
)
pending_tasks = {}
def should_continue(site: dict, page_results: List[Any]) -> bool:
"""按既有站点分页规则判断是否提交下一页资源请求。"""
search_keyword = (
mediainfo.imdb_id
if area == "imdbid" and mediainfo
else keyword
)
return self._should_continue_search_pages(
site=site,
page_results=page_results,
keyword=search_keyword,
)
def submit_site_page(site: dict, page_index: int):
"""
提交异步站点页搜索任务,并记录该任务对应的站点和页码位置。
"""
search_page = search_pages[page_index]
search_keyword = mediainfo.imdb_id if area == "imdbid" and mediainfo else keyword
task = asyncio.create_task(search_site_page(site=site, search_page=search_page))
pending_tasks[task] = (site, page_index, search_page, search_keyword)
for site in indexer_sites:
submit_site_page(site=site, page_index=0)
try:
while pending_tasks:
if global_vars.is_system_stopped:
break
done_tasks, _ = await asyncio.wait(
pending_tasks.keys(),
return_when=asyncio.FIRST_COMPLETED,
page_iterator = self._iter_site_page_results(
indexer_sites=indexer_sites,
search_pages=search_pages,
search_page=search_site_page,
should_continue=should_continue,
task_owner="chain.search.media.site_page",
)
async with aclosing(page_iterator):
async for site, search_page, result, continued in page_iterator:
finish_count += 1
results.extend(result)
if not continued:
logger.debug(
f"{site.get('name')}{search_page} 页返回 {len(result)} 条,停止继续翻页"
)
logger.info(f"站点搜索进度:{finish_count} / {total_num}")
await progress.update(
value=finish_count / total_num * 100,
text=(
f"正在搜索{keyword or ''},已完成 "
f"{finish_count} / {total_num} 个请求 ..."
),
)
for future in done_tasks:
site, page_index, search_page, search_keyword = pending_tasks.pop(future)
finish_count += 1
result = await future
if result:
results.extend(result)
if (
self._should_continue_search_pages(
site=site, page_results=result, keyword=search_keyword
)
and page_index + 1 < len(search_pages)
):
submit_site_page(site=site, page_index=page_index + 1)
else:
logger.debug(
f"{site.get('name')}{search_page} 页返回 {len(result or [])} 条,停止继续翻页"
)
logger.info(f"站点搜索进度:{finish_count} / {total_num}")
await progress.update(value=finish_count / total_num * 100,
text=f"正在搜索{keyword or ''},已完成 {finish_count} / {total_num} 个请求 ...")
finally:
for task in pending_tasks:
if not task.done():
task.cancel()
if pending_tasks:
await asyncio.gather(*pending_tasks.keys(), return_exceptions=True)
# 计算耗时
end_time = datetime.now()
@@ -2630,89 +2675,69 @@ class SearchChain(ChainBase):
"total": total_num
}
semaphore = asyncio.Semaphore(
self.runtime_config.search_threadpool_size or total_num
)
async def search_site(site: dict, search_page: int) -> List[TorrentInfo]:
"""
搜索单个站点页,用于渐进式返回入口。
"""
async with semaphore:
if area == "imdbid":
site_result = await self.async_search_site_torrents(site=site,
keyword=mediainfo.imdb_id if mediainfo else None,
mtype=mediainfo.type if mediainfo else mtype,
page=search_page)
else:
site_result = await self.async_search_site_torrents(site=site,
keyword=keyword,
mtype=mediainfo.type if mediainfo else mtype,
page=search_page)
return site_result or []
"""调用既有站点资源接口,具体并发由统一逐页编排器控制。"""
search_keyword = (
mediainfo.imdb_id
if area == "imdbid" and mediainfo
else keyword
)
return await self.async_search_site_torrents(
site=site,
keyword=search_keyword,
mtype=mediainfo.type if mediainfo else mtype,
page=search_page,
)
tasks = {}
def submit_site_page(site: dict, page_index: int):
"""
提交渐进式站点页搜索任务,并保留站点和页码上下文。
"""
search_page = search_pages[page_index]
search_keyword = mediainfo.imdb_id if area == "imdbid" and mediainfo else keyword
task = asyncio.create_task(search_site(site=site, search_page=search_page))
tasks[task] = (site, page_index, search_page, search_keyword)
for site in indexer_sites:
submit_site_page(site=site, page_index=0)
def should_continue(site: dict, page_results: List[Any]) -> bool:
"""按既有站点分页规则判断是否提交下一页资源请求。"""
search_keyword = (
mediainfo.imdb_id
if area == "imdbid" and mediainfo
else keyword
)
return self._should_continue_search_pages(
site=site,
page_results=page_results,
keyword=search_keyword,
)
results_count = len(plugin_results)
try:
while tasks:
if global_vars.is_system_stopped:
break
done_tasks, _ = await asyncio.wait(
tasks.keys(),
return_when=asyncio.FIRST_COMPLETED,
page_iterator = self._iter_site_page_results(
indexer_sites=indexer_sites,
search_pages=search_pages,
search_page=search_site,
should_continue=should_continue,
task_owner="chain.search.media.site_page",
)
async with aclosing(page_iterator):
async for site, search_page, result, continued in page_iterator:
finish_count += 1
results_count += len(result)
if not continued:
logger.debug(
f"{site.get('name')}{search_page} 页返回 {len(result)} 条,停止继续翻页"
)
logger.info(f"站点搜索进度:{finish_count} / {total_num}")
progress_value = finish_count / total_num * 100
progress_text = (
f"正在搜索{keyword or ''},已完成 "
f"{finish_count} / {total_num} 个请求 ..."
)
for future in done_tasks:
site, page_index, search_page, search_keyword = tasks.pop(future)
finish_count += 1
result = await future
results_count += len(result)
if (
self._should_continue_search_pages(
site=site, page_results=result, keyword=search_keyword
)
and page_index + 1 < len(search_pages)
):
submit_site_page(site=site, page_index=page_index + 1)
else:
logger.debug(
f"{site.get('name')}{search_page} 页返回 {len(result)} 条,停止继续翻页"
)
logger.info(f"站点搜索进度:{finish_count} / {total_num}")
progress_value = finish_count / total_num * 100
progress_text = f"正在搜索{keyword or ''},已完成 {finish_count} / {total_num} 个请求 ..."
await progress.update(value=progress_value, text=progress_text)
yield {
"type": "append",
"stage": "searching",
"value": progress_value,
"text": progress_text,
"items": result,
"site": site.get("name"),
"site_id": site.get("id"),
"page": search_page,
"finished": finish_count,
"total": total_num,
"total_items": results_count
}
finally:
for task in tasks:
if not task.done():
task.cancel()
if tasks:
await asyncio.gather(*tasks.keys(), return_exceptions=True)
await progress.update(value=progress_value, text=progress_text)
yield {
"type": "append",
"stage": "searching",
"value": progress_value,
"text": progress_text,
"items": result,
"site": site.get("name"),
"site_id": site.get("id"),
"page": search_page,
"finished": finish_count,
"total": total_num,
"total_items": results_count
}
end_time = datetime.now()
await progress.update(value=100,
@@ -2754,65 +2779,46 @@ class SearchChain(ChainBase):
await progress.update(value=0,
text=f"开始搜索字幕,共 {len(indexer_sites)} 个站点,{len(search_pages)} 页 ...")
results = []
semaphore = asyncio.Semaphore(
self.runtime_config.search_threadpool_size or total_num
)
async def search_site_page(site: dict, search_page: int) -> List[SubtitleInfo]:
"""
控制单次字幕站点页请求的并发量,并返回该页的字幕列表。
"""
async with semaphore:
return await self.async_search_subtitles(
site=site, keyword=keyword, page=search_page
"""调用既有字幕接口,具体并发由统一逐页编排器控制。"""
return await self.async_search_subtitles(
site=site,
keyword=keyword,
page=search_page,
)
def should_continue(site: dict, page_results: List[Any]) -> bool:
"""按既有字幕分页规则判断是否提交下一页请求。"""
return self._should_continue_subtitle_search_pages(
site=site,
page_results=page_results,
)
page_iterator = self._iter_site_page_results(
indexer_sites=indexer_sites,
search_pages=search_pages,
search_page=search_site_page,
should_continue=should_continue,
task_owner="chain.search.subtitle.site_page",
)
async with aclosing(page_iterator):
async for site, search_page, result, continued in page_iterator:
finish_count += 1
results.extend(result)
if not continued:
logger.debug(
f"{site.get('name')} 字幕第 {search_page} 页返回 {len(result)} 条,停止继续翻页"
)
logger.info(f"站点字幕搜索进度:{finish_count} / {total_num}")
await progress.update(
value=finish_count / total_num * 100,
text=(
f"正在搜索字幕{keyword or ''},已完成 "
f"{finish_count} / {total_num} 个请求 ..."
),
)
pending_tasks = {}
def submit_site_page(site: dict, page_index: int):
"""
提交异步字幕站点页搜索任务,并记录站点和页码位置。
"""
search_page = search_pages[page_index]
task = asyncio.create_task(search_site_page(site=site, search_page=search_page))
pending_tasks[task] = (site, page_index, search_page)
for site in indexer_sites:
submit_site_page(site=site, page_index=0)
try:
while pending_tasks:
if global_vars.is_system_stopped:
break
done_tasks, _ = await asyncio.wait(
pending_tasks.keys(),
return_when=asyncio.FIRST_COMPLETED,
)
for future in done_tasks:
site, page_index, search_page = pending_tasks.pop(future)
finish_count += 1
result = await future
if result:
results.extend(result)
if (
self._should_continue_subtitle_search_pages(site=site, page_results=result)
and page_index + 1 < len(search_pages)
):
submit_site_page(site=site, page_index=page_index + 1)
else:
logger.debug(
f"{site.get('name')} 字幕第 {search_page} 页返回 {len(result or [])} 条,停止继续翻页"
)
logger.info(f"站点字幕搜索进度:{finish_count} / {total_num}")
await progress.update(value=finish_count / total_num * 100,
text=f"正在搜索字幕{keyword or ''},已完成 {finish_count} / {total_num} 个请求 ...")
finally:
for task in pending_tasks:
if not task.done():
task.cancel()
if pending_tasks:
await asyncio.gather(*pending_tasks.keys(), return_exceptions=True)
end_time = datetime.now()
await progress.update(value=100,
text=f"站点字幕搜索完成,有效字幕数:{len(results)},总耗时 {(end_time - start_time).seconds}")
@@ -2871,79 +2877,57 @@ class SearchChain(ChainBase):
"total": total_num
}
semaphore = asyncio.Semaphore(
self.runtime_config.search_threadpool_size or total_num
)
async def search_site(site: dict, search_page: int) -> List[SubtitleInfo]:
"""
搜索单个站点字幕页,用于渐进式返回入口。
"""
async with semaphore:
site_result = await self.async_search_subtitles(
site=site, keyword=keyword, page=search_page
)
return site_result or []
"""调用既有字幕接口,具体并发由统一逐页编排器控制。"""
return await self.async_search_subtitles(
site=site,
keyword=keyword,
page=search_page,
)
tasks = {}
def submit_site_page(site: dict, page_index: int):
"""
提交渐进式字幕站点页搜索任务,并保留站点和页码上下文。
"""
search_page = search_pages[page_index]
task = asyncio.create_task(search_site(site=site, search_page=search_page))
tasks[task] = (site, page_index, search_page)
for site in indexer_sites:
submit_site_page(site=site, page_index=0)
def should_continue(site: dict, page_results: List[Any]) -> bool:
"""按既有字幕分页规则判断是否提交下一页请求。"""
return self._should_continue_subtitle_search_pages(
site=site,
page_results=page_results,
)
results_count = 0
try:
while tasks:
if global_vars.is_system_stopped:
break
done_tasks, _ = await asyncio.wait(
tasks.keys(),
return_when=asyncio.FIRST_COMPLETED,
page_iterator = self._iter_site_page_results(
indexer_sites=indexer_sites,
search_pages=search_pages,
search_page=search_site,
should_continue=should_continue,
task_owner="chain.search.subtitle.site_page",
)
async with aclosing(page_iterator):
async for site, search_page, result, continued in page_iterator:
finish_count += 1
results_count += len(result)
if not continued:
logger.debug(
f"{site.get('name')} 字幕第 {search_page} 页返回 {len(result)} 条,停止继续翻页"
)
logger.info(f"站点字幕搜索进度:{finish_count} / {total_num}")
progress_value = finish_count / total_num * 100
progress_text = (
f"正在搜索字幕{keyword or ''},已完成 "
f"{finish_count} / {total_num} 个请求 ..."
)
for future in done_tasks:
site, page_index, search_page = tasks.pop(future)
finish_count += 1
result = await future
results_count += len(result)
if (
self._should_continue_subtitle_search_pages(site=site, page_results=result)
and page_index + 1 < len(search_pages)
):
submit_site_page(site=site, page_index=page_index + 1)
else:
logger.debug(
f"{site.get('name')} 字幕第 {search_page} 页返回 {len(result)} 条,停止继续翻页"
)
logger.info(f"站点字幕搜索进度:{finish_count} / {total_num}")
progress_value = finish_count / total_num * 100
progress_text = f"正在搜索字幕{keyword or ''},已完成 {finish_count} / {total_num} 个请求 ..."
await progress.update(value=progress_value, text=progress_text)
yield {
"type": "append",
"stage": "searching",
"value": progress_value,
"text": progress_text,
"items": result,
"site": site.get("name"),
"site_id": site.get("id"),
"page": search_page,
"finished": finish_count,
"total": total_num,
"total_items": results_count
}
finally:
for task in tasks:
if not task.done():
task.cancel()
if tasks:
await asyncio.gather(*tasks.keys(), return_exceptions=True)
await progress.update(value=progress_value, text=progress_text)
yield {
"type": "append",
"stage": "searching",
"value": progress_value,
"text": progress_text,
"items": result,
"site": site.get("name"),
"site_id": site.get("id"),
"page": search_page,
"finished": finish_count,
"total": total_num,
"total_items": results_count
}
end_time = datetime.now()
await progress.update(value=100,