diff --git a/tests/builtin_plugins/plugin_store/test_download_strategy.py b/tests/builtin_plugins/plugin_store/test_download_strategy.py new file mode 100644 index 00000000..52e2592c --- /dev/null +++ b/tests/builtin_plugins/plugin_store/test_download_strategy.py @@ -0,0 +1,437 @@ +import asyncio +from pathlib import Path +import shutil + +import pytest +from pytest_mock import MockerFixture + + +async def _run_git(*args: str) -> None: + process = await asyncio.create_subprocess_exec( + "git", + *args, + stdout=asyncio.subprocess.PIPE, + stderr=asyncio.subprocess.PIPE, + ) + _, stderr = await process.communicate() + assert process.returncode == 0, stderr.decode(errors="replace") + + +def _plugin_info(*, ali_url: str | None = None): + from zhenxun.builtin_plugins.plugin_store.models import StorePluginInfo + from zhenxun.utils.enum import PluginType + + return StorePluginInfo( + name="测试插件", + module="demo", + module_path="demo", + description="", + usage="", + author="tester", + version="1.0.0", + plugin_type=PluginType.NORMAL, + is_dir=True, + github_url="https://github.com/example/demo", + ali_url=ali_url, + ) + + +def test_source_order() -> None: + from zhenxun.builtin_plugins.plugin_store.data_source import StoreManager + from zhenxun.utils.repo_utils.models import RepoType + + assert StoreManager._get_source_order(None) == ( + RepoType.ALIYUN, + RepoType.GITHUB, + ) + assert StoreManager._get_source_order("ali") == (RepoType.ALIYUN,) + assert StoreManager._get_source_order("git") == (RepoType.GITHUB,) + + +def test_repository_branch_is_resolved_per_source() -> None: + from zhenxun.builtin_plugins.plugin_store.data_source import StoreManager + from zhenxun.utils.repo_utils.models import RepoType + + plugin_info = _plugin_info() + plugin_info.github_url = "https://github.com/example/demo/tree/master" + + assert ( + StoreManager._get_plugin_repository_branch(plugin_info, RepoType.ALIYUN, "main") + == "main" + ) + assert ( + StoreManager._get_plugin_repository_branch(plugin_info, RepoType.GITHUB, "main") + == "master" + ) + + +@pytest.mark.parametrize("is_external", [False, True]) +async def test_default_source_falls_back_to_github_for_all_plugins( + mocker: MockerFixture, + tmp_path: Path, + is_external: bool, +) -> None: + from zhenxun.builtin_plugins.plugin_store.data_source import StoreManager + from zhenxun.utils.repo_utils.models import ( + FileDownloadResult, + RepoFileInfo, + RepoType, + ) + + mock_base_path = mocker.patch( + "zhenxun.builtin_plugins.plugin_store.data_source.BASE_PATH", + new=tmp_path / "zhenxun", + ) + source_calls: list[RepoType] = [] + download_calls: list[tuple[RepoType, list[str]]] = [] + files = [ + RepoFileInfo(path="demo/__init__.py", is_dir=False), + RepoFileInfo(path="demo/assets/icon.png", is_dir=False), + RepoFileInfo(path="demo/requirements.txt", is_dir=False), + ] + + async def list_directory_files( + repo_url: str, + directory_path: str, + branch: str, + repo_type: RepoType, + ) -> list[RepoFileInfo]: + source_calls.append(repo_type) + if repo_type == RepoType.ALIYUN: + raise RuntimeError("aliyun unavailable") + return files + + async def download_files( + repo_url: str, + file_path: list[tuple[str, Path]], + branch: str, + repo_type: RepoType, + ignore_error: bool = False, + ) -> FileDownloadResult: + download_calls.append((repo_type, [path for path, _ in file_path])) + for source_path, destination_path in file_path: + destination_path.parent.mkdir(parents=True, exist_ok=True) + if source_path.endswith(".png"): + destination_path.write_bytes(b"\x89PNG\r\n\x1a\n") + else: + destination_path.write_text(source_path, encoding="utf-8") + return FileDownloadResult( + repo_type=repo_type, + repo_name="demo", + file_path=file_path, + version=branch, + success=True, + ) + + mocker.patch( + "zhenxun.builtin_plugins.plugin_store.data_source." + "RepoFileManager.list_directory_files", + side_effect=list_directory_files, + ) + mocker.patch( + "zhenxun.builtin_plugins.plugin_store.data_source." + "RepoFileManager.download_files", + side_effect=download_files, + ) + install_requirement = mocker.patch( + "zhenxun.builtin_plugins.plugin_store.data_source." + "VirtualEnvPackageManager.install_requirement", + ) + + await StoreManager.install_plugin_with_repo( + _plugin_info(), + is_external=is_external, + ) + + assert source_calls == [RepoType.ALIYUN, RepoType.GITHUB] + assert download_calls == [ + ( + RepoType.GITHUB, + [ + "demo/__init__.py", + "demo/assets/icon.png", + "demo/requirements.txt", + ], + ) + ] + assert ( + mock_base_path / "plugins" / "demo" / "assets" / "icon.png" + ).read_bytes() == b"\x89PNG\r\n\x1a\n" + install_requirement.assert_awaited_once() + + +async def test_zero_byte_aliyun_binary_falls_back_to_github( + mocker: MockerFixture, + tmp_path: Path, +) -> None: + from zhenxun.builtin_plugins.plugin_store.data_source import StoreManager + from zhenxun.utils.repo_utils.models import ( + FileDownloadResult, + RepoFileInfo, + RepoType, + ) + + mock_base_path = mocker.patch( + "zhenxun.builtin_plugins.plugin_store.data_source.BASE_PATH", + new=tmp_path / "zhenxun", + ) + calls: list[RepoType] = [] + files = [ + RepoFileInfo(path="demo/__init__.py", is_dir=False), + RepoFileInfo(path="demo/icon.png", is_dir=False), + ] + + mocker.patch( + "zhenxun.builtin_plugins.plugin_store.data_source." + "RepoFileManager.list_directory_files", + return_value=files, + ) + + async def download_files( + repo_url: str, + file_path: list[tuple[str, Path]], + branch: str, + repo_type: RepoType, + ignore_error: bool = False, + ) -> FileDownloadResult: + calls.append(repo_type) + for source_path, destination_path in file_path: + destination_path.parent.mkdir(parents=True, exist_ok=True) + if source_path.endswith(".png"): + content = b"" if repo_type == RepoType.ALIYUN else b"image" + destination_path.write_bytes(content) + else: + destination_path.write_text("", encoding="utf-8") + return FileDownloadResult( + repo_type=repo_type, + repo_name="demo", + file_path=file_path, + version=branch, + success=True, + ) + + mocker.patch( + "zhenxun.builtin_plugins.plugin_store.data_source." + "RepoFileManager.download_files", + side_effect=download_files, + ) + mocker.patch( + "zhenxun.builtin_plugins.plugin_store.data_source." + "VirtualEnvPackageManager.install_requirement", + ) + + await StoreManager.install_plugin_with_repo(_plugin_info(), is_external=True) + + assert calls[:2] == [RepoType.ALIYUN, RepoType.GITHUB] + assert (mock_base_path / "plugins" / "demo" / "icon.png").read_bytes() == b"image" + + +async def test_forced_aliyun_does_not_fall_back( + mocker: MockerFixture, + tmp_path: Path, +) -> None: + from zhenxun.builtin_plugins.plugin_store.data_source import StoreManager + from zhenxun.builtin_plugins.plugin_store.exceptions import PluginStoreException + from zhenxun.utils.repo_utils.models import RepoType + + mocker.patch( + "zhenxun.builtin_plugins.plugin_store.data_source.BASE_PATH", + new=tmp_path / "zhenxun", + ) + list_directory_files = mocker.patch( + "zhenxun.builtin_plugins.plugin_store.data_source." + "RepoFileManager.list_directory_files", + side_effect=RuntimeError("aliyun unavailable"), + ) + + with pytest.raises(PluginStoreException, match="阿里云"): + await StoreManager.install_plugin_with_repo( + _plugin_info(), + source="ali", + ) + + assert list_directory_files.await_count == 1 + await_args = list_directory_files.await_args + assert await_args is not None + assert await_args.kwargs["repo_type"] == RepoType.ALIYUN + + +async def test_root_plugin_uses_exact_sparse_paths( + mocker: MockerFixture, + tmp_path: Path, +) -> None: + from zhenxun.builtin_plugins.plugin_store.data_source import StoreManager + from zhenxun.utils.repo_utils.models import ( + FileDownloadResult, + RepoFileInfo, + RepoType, + ) + + mock_base_path = mocker.patch( + "zhenxun.builtin_plugins.plugin_store.data_source.BASE_PATH", + new=tmp_path / "zhenxun", + ) + plugin_info = _plugin_info( + ali_url="https://codeup.aliyun.com/organization/group/demo-mirror" + ) + plugin_info.module_path = "." + files = [ + RepoFileInfo(path="__init__.py", is_dir=False), + RepoFileInfo(path="assets/icon.png", is_dir=False), + RepoFileInfo(path="requirements.txt", is_dir=False), + ] + mocker.patch( + "zhenxun.builtin_plugins.plugin_store.data_source." + "RepoFileManager.list_directory_files", + return_value=files, + ) + downloaded_paths: list[str] = [] + + async def download_files( + repo_url: str, + file_path: list[tuple[str, Path]], + branch: str, + repo_type: RepoType, + ignore_error: bool = False, + ) -> FileDownloadResult: + downloaded_paths.extend(path for path, _ in file_path) + for source_path, destination_path in file_path: + destination_path.parent.mkdir(parents=True, exist_ok=True) + content = b"image" if source_path.endswith(".png") else b"" + destination_path.write_bytes(content) + return FileDownloadResult( + repo_type=repo_type, + repo_name="demo-mirror", + file_path=file_path, + version=branch, + success=True, + ) + + mocker.patch( + "zhenxun.builtin_plugins.plugin_store.data_source." + "RepoFileManager.download_files", + side_effect=download_files, + ) + mocker.patch( + "zhenxun.builtin_plugins.plugin_store.data_source." + "VirtualEnvPackageManager.install_requirement", + ) + + await StoreManager.install_plugin_with_repo(plugin_info, source="ali") + + assert downloaded_paths == [ + "__init__.py", + "assets/icon.png", + "requirements.txt", + ] + assert ( + mock_base_path / "plugins" / "demo" / "assets" / "icon.png" + ).read_bytes() == b"image" + + +async def test_repo_manager_sparse_checkout_preserves_exact_paths( + mocker: MockerFixture, + tmp_path: Path, +) -> None: + from zhenxun.utils.repo_utils import RepoFileManager + from zhenxun.utils.repo_utils.models import RepoType + + async def sparse_checkout_clone( + repo_url: str, + branch: str, + sparse_path: list[str], + target_dir: Path, + ) -> list[str]: + assert repo_url == "https://github.com/example/demo.git" + assert branch == "master" + assert sparse_path == ["demo/assets/icon.png"] + source = target_dir / sparse_path[0] + source.parent.mkdir(parents=True, exist_ok=True) + source.write_bytes(b"image") + return sparse_path + + sparse_checkout = mocker.patch( + "zhenxun.utils.repo_utils.file_manager.sparse_checkout_clone", + side_effect=sparse_checkout_clone, + ) + target = tmp_path / "target" / "icon.png" + result = await RepoFileManager.download_files( + "https://github.com/example/demo/tree/master", + [("demo/assets/icon.png", target)], + "master", + repo_type=RepoType.GITHUB, + ) + + assert result.success + assert target.read_bytes() == b"image" + sparse_checkout.assert_awaited_once() + + +async def test_sparse_checkout_retries_git_fetch( + mocker: MockerFixture, + tmp_path: Path, +) -> None: + from zhenxun.utils.repo_utils.utils import sparse_checkout_clone + + mocker.patch("zhenxun.utils.repo_utils.utils.check_git", return_value=True) + sleep = mocker.patch("zhenxun.utils.repo_utils.utils.asyncio.sleep") + fetch_attempts = 0 + fetch_timeouts: list[float | None] = [] + + async def run_git_command( + command: str | list[str], + cwd: Path | None = None, + timeout_seconds: float | None = None, + ) -> tuple[bool, str, str]: + nonlocal fetch_attempts + if isinstance(command, list) and "fetch" in command: + fetch_attempts += 1 + fetch_timeouts.append(timeout_seconds) + if fetch_attempts < 3: + return False, "", "connection reset" + return True, "", "" + + mocker.patch( + "zhenxun.utils.repo_utils.utils.run_git_command", + side_effect=run_git_command, + ) + + downloaded = await sparse_checkout_clone( + repo_url="https://github.com/example/demo", + branch="main", + sparse_path=["demo/__init__.py"], + target_dir=tmp_path / "target", + ) + + assert downloaded == [] + assert fetch_attempts == 3 + assert fetch_timeouts == [60, 60, 60] + assert sleep.await_count == 2 + + +@pytest.mark.skipif(shutil.which("git") is None, reason="git is not installed") +async def test_git_checkout_preserves_binary(tmp_path: Path) -> None: + from zhenxun.utils.repo_utils.utils import sparse_checkout_clone + + source_repo = tmp_path / "source" + source_repo.mkdir() + await _run_git("init", "-b", "main", str(source_repo)) + await _run_git("-C", str(source_repo), "config", "user.name", "test") + await _run_git("-C", str(source_repo), "config", "user.email", "test@example.com") + binary_content = b"\x89PNG\r\n\x1a\n\x00\x01\xffbinary" + (source_repo / "icon.png").write_bytes(binary_content) + (source_repo / "__init__.py").write_text("", encoding="utf-8") + await _run_git("-C", str(source_repo), "add", ".") + await _run_git("-C", str(source_repo), "commit", "-m", "test") + + target_dir = tmp_path / "target" + downloaded = await sparse_checkout_clone( + repo_url=source_repo.as_uri(), + branch="main", + sparse_path=["icon.png", "__init__.py"], + target_dir=target_dir, + ) + + assert downloaded == ["icon.png", "__init__.py"] + assert (target_dir / "icon.png").read_bytes() == binary_content + assert not (target_dir / ".git").exists() diff --git a/zhenxun/builtin_plugins/plugin_store/data_source.py b/zhenxun/builtin_plugins/plugin_store/data_source.py index 4278fabe..7bbbc612 100644 --- a/zhenxun/builtin_plugins/plugin_store/data_source.py +++ b/zhenxun/builtin_plugins/plugin_store/data_source.py @@ -1,12 +1,12 @@ import os from pathlib import Path -import random import shutil +import tempfile +from typing import ClassVar import ujson as json from zhenxun.builtin_plugins.plugin_store.models import StorePluginInfo -from zhenxun.configs.path_config import TEMP_PATH from zhenxun.models.plugin_info import PluginInfo from zhenxun.services.cache.bounded_ttl import BoundedTTLCache from zhenxun.services.log import logger @@ -52,6 +52,58 @@ def row_style(column: str, text: str) -> RowStyle: class StoreManager: + _SOURCE_NAMES: ClassVar[dict[RepoType, str]] = { + RepoType.ALIYUN: "阿里云", + RepoType.GITHUB: "GitHub", + } + _BINARY_EXTENSIONS: ClassVar[frozenset[str]] = frozenset( + { + ".7z", + ".avi", + ".bin", + ".bmp", + ".class", + ".dat", + ".db", + ".dll", + ".doc", + ".docx", + ".dylib", + ".eot", + ".exe", + ".flv", + ".gif", + ".gz", + ".ico", + ".jpeg", + ".jpg", + ".mov", + ".mp3", + ".mp4", + ".otf", + ".pdf", + ".png", + ".ppt", + ".pptx", + ".pyc", + ".rar", + ".so", + ".svg", + ".tar", + ".tif", + ".tiff", + ".ttf", + ".webp", + ".wmv", + ".woff", + ".woff2", + ".xls", + ".xlsx", + ".xz", + ".zip", + } + ) + @classmethod def _resolve_local_plugin_path( cls, plugin_info: StorePluginInfo, *, is_external: bool @@ -333,87 +385,214 @@ class StoreManager: 参数: plugin_info: 插件信息 - is_external: 是否是外部仓库 - source: 源 + is_external: 是否是外部仓库(保留用于兼容旧调用) + source: 强制使用的源,ali 为阿里云,git 为 GitHub; + 不指定时优先阿里云,失败后回退 GitHub """ - repo_type = RepoType.GITHUB if is_external else None - if not is_external: - repo_type = RepoType.ALIYUN - elif (source is None and plugin_info.ali_url) or source == "ali": - repo_type = RepoType.ALIYUN - elif source == "git": - repo_type = RepoType.GITHUB + source_order = cls._get_source_order(source) + errors: list[str] = [] + + with tempfile.TemporaryDirectory(prefix="zhenxun_plugin_store_") as temp_dir: + staged_result: tuple[list[tuple[Path, Path]], list[Path]] | None = None + selected_source: RepoType | None = None + + for repo_type in source_order: + source_name = cls._SOURCE_NAMES[repo_type] + staging_root = Path(temp_dir) / repo_type.value + try: + staged_result = await cls._download_plugin_to_staging( + plugin_info, + repo_type, + branch, + staging_root, + ) + selected_source = repo_type + logger.info( + f"插件 {plugin_info.name} 使用{source_name}下载成功", + LOG_COMMAND, + ) + break + except Exception as e: + errors.append(f"{source_name}: {e}") + if repo_type != source_order[-1]: + logger.warning( + f"插件 {plugin_info.name} 使用{source_name}下载失败," + "尝试 GitHub", + LOG_COMMAND, + e=e, + ) + + if staged_result is None or selected_source is None: + raise PluginStoreException( + f"插件 {plugin_info.name} 下载失败({';'.join(errors)})" + ) + + deploy_files, requirement_files = staged_result + for requirement_file in requirement_files: + logger.info( + f"开始安装插件 {plugin_info.module_path} " + f"依赖文件: {requirement_file}", + LOG_COMMAND, + ) + await VirtualEnvPackageManager.install_requirement(requirement_file) + + for staged_path, destination_path in deploy_files: + destination_path.parent.mkdir(parents=True, exist_ok=True) + shutil.copy2(staged_path, destination_path) + + @staticmethod + def _get_source_order(source: str | None) -> tuple[RepoType, ...]: + """解析插件下载源。""" + if source is None: + return (RepoType.ALIYUN, RepoType.GITHUB) + if source == "ali": + return (RepoType.ALIYUN,) + if source == "git": + return (RepoType.GITHUB,) + raise PluginStoreException(f"源类型错误: {source},请使用 ali 或 git") + + @staticmethod + def _get_repository_url(plugin_info: StorePluginInfo, repo_type: RepoType) -> str: + """获取指定下载源对应的仓库地址。""" + if repo_type == RepoType.ALIYUN and plugin_info.ali_url: + return plugin_info.ali_url + if plugin_info.github_url: + return plugin_info.github_url + raise PluginStoreException(f"插件 {plugin_info.name} 缺少仓库地址") + + @staticmethod + def _get_repository_branch(repo_url: str | None, default_branch: str) -> str: + """优先使用仓库 URL 中显式指定的分支、标签或提交。""" + if repo_url and "/tree/" in repo_url: + _, _, ref = repo_url.partition("/tree/") + if ref := ref.strip("/"): + return ref + return default_branch + + @classmethod + def _get_plugin_repository_branch( + cls, + plugin_info: StorePluginInfo, + repo_type: RepoType, + default_branch: str, + ) -> str: + """按下载源独立解析分支,避免把 GitHub 分支套到阿里云镜像。""" + branch_source_url = ( + plugin_info.ali_url + if repo_type == RepoType.ALIYUN + else plugin_info.github_url + ) + return cls._get_repository_branch(branch_source_url, default_branch) + + @classmethod + async def _download_plugin_to_staging( + cls, + plugin_info: StorePluginInfo, + repo_type: RepoType, + default_branch: str, + staging_root: Path, + ) -> tuple[list[tuple[Path, Path]], list[Path]]: + """从单一仓库源完整下载插件到临时目录。""" + repo_url = cls._get_repository_url(plugin_info, repo_type) + branch = cls._get_plugin_repository_branch( + plugin_info, + repo_type, + default_branch, + ) module_path = plugin_info.module_path - is_dir = plugin_info.is_dir - github_url = plugin_info.github_url - assert github_url - replace_module_path = module_path.replace(".", "/").lstrip("/") - plugin_module = plugin_info.module - if is_dir: + repository_plugin_path = module_path.replace(".", "/").strip("/") + + if plugin_info.is_dir: files = await RepoFileManager.list_directory_files( - github_url, replace_module_path, branch, repo_type=repo_type + repo_url, + repository_plugin_path, + branch, + repo_type=repo_type, ) else: - files = [RepoFileInfo(path=f"{replace_module_path}.py", is_dir=False)] - if is_dir: - target_dir = BASE_PATH / "plugins" / plugin_module - else: - target_dir = BASE_PATH / "plugins" + if not repository_plugin_path: + raise PluginStoreException( + f"插件 {plugin_info.name} 的模块路径不能为空" + ) + files = [RepoFileInfo(path=f"{repository_plugin_path}.py", is_dir=False)] + files = [file for file in files if not file.is_dir] + if not files: + raise PluginStoreException( + f"仓库中未找到插件目录: {plugin_info.module_path}" + ) + + target_root = ( + BASE_PATH / "plugins" / plugin_info.module + if plugin_info.is_dir + else BASE_PATH / "plugins" + ) download_files: list[tuple[str, Path]] = [] + deploy_files: list[tuple[Path, Path]] = [] for file in files: - src_path = file.path - if is_dir: - dst_path = target_dir / Path(src_path).relative_to(replace_module_path) - else: - dst_path = target_dir / f"{plugin_module}.py" + source_path = Path(file.path) + if source_path.is_absolute() or ".." in source_path.parts: + raise PluginStoreException(f"仓库包含不安全的文件路径: {file.path}") + + staged_path = staging_root / source_path + if plugin_info.is_dir: + plugin_root = ( + Path(repository_plugin_path) if repository_plugin_path else Path() + ) + try: + relative_path = source_path.relative_to(plugin_root) + except ValueError as e: + raise PluginStoreException( + f"插件文件不在模块目录内: {file.path}" + ) from e + destination_path = target_root / relative_path + else: + destination_path = target_root / f"{plugin_info.module}.py" + + download_files.append((file.path, staged_path)) + deploy_files.append((staged_path, destination_path)) + + required_download_files = download_files.copy() + requirement_files = [ + staging_root / Path(file.path) + for file in files + if Path(file.path).name in {"requirement.txt", "requirements.txt"} + ] + root_requirements: list[tuple[str, Path]] = [] + if not requirement_files: + root_requirements = [ + ("requirement.txt", staging_root / "requirement.txt"), + ("requirements.txt", staging_root / "requirements.txt"), + ] + download_files.extend(root_requirements) - download_files.append((src_path, dst_path)) - rand = random.randint(1, 10000) - requirement_path_ = TEMP_PATH / f"plugin_store_{rand}_req.txt" - requirements_path_ = TEMP_PATH / f"plugin_store_{rand}_reqs.txt" - download_files.append(("requirement.txt", requirement_path_)) - download_files.append(("requirements.txt", requirements_path_)) result = await RepoFileManager.download_files( - github_url, + repo_url, download_files, branch, repo_type=repo_type, - ignore_error=True, + ignore_error=bool(root_requirements), ) if not result.success: - raise PluginStoreException(result.error_message) + raise PluginStoreException(result.error_message or "未知下载错误") - requirement_paths = [ - file - for file in files - if file.path.endswith("requirement.txt") - or file.path.endswith("requirements.txt") - ] + for source_path, staged_path in required_download_files: + if not staged_path.is_file(): + raise PluginStoreException(f"插件文件下载不完整: {source_path}") + if ( + Path(source_path).suffix.lower() in cls._BINARY_EXTENSIONS + and staged_path.stat().st_size == 0 + ): + raise PluginStoreException(f"二进制文件下载为空: {source_path}") - is_install_req = False - for requirement_path in requirement_paths: - requirement_file = target_dir / Path(requirement_path.path).relative_to( - replace_module_path - ) - if requirement_file.exists(): - is_install_req = True - await VirtualEnvPackageManager.install_requirement(requirement_file) + requirement_files = [path for path in requirement_files if path.is_file()] + if root_requirements: + requirement_files = [ + path for _, path in root_requirements if path.is_file() + ] - if not is_install_req: - if requirement_path_.exists(): - logger.info( - f"开始安装插件 {module_path} 依赖文件: {requirement_path_}", - LOG_COMMAND, - ) - await VirtualEnvPackageManager.install_requirement(requirement_path_) - if requirements_path_.exists(): - logger.info( - f"开始安装插件 {module_path} 依赖文件: {requirements_path_}", - LOG_COMMAND, - ) - await VirtualEnvPackageManager.install_requirement(requirements_path_) + return deploy_files, requirement_files @classmethod async def remove_plugin(cls, index_or_module: str) -> str: diff --git a/zhenxun/utils/repo_utils/utils.py b/zhenxun/utils/repo_utils/utils.py index 4205d6d5..023dc720 100644 --- a/zhenxun/utils/repo_utils/utils.py +++ b/zhenxun/utils/repo_utils/utils.py @@ -4,11 +4,14 @@ import asyncio import base64 +import contextlib from pathlib import Path import re import shutil import tempfile +import psutil + from zhenxun.services.log import logger from .config import LOG_COMMAND, RepoConfig @@ -48,7 +51,9 @@ async def clean_git(cwd: Path): async def run_git_command( - command: str | list[str], cwd: Path | None = None + command: str | list[str], + cwd: Path | None = None, + timeout_seconds: float | None = None, ) -> tuple[bool, str, str]: """ 运行git命令,实时输出 stderr 进度信息(如 git clone --progress)。 @@ -56,6 +61,7 @@ async def run_git_command( 参数: command: 命令字符串或参数列表 cwd: 工作目录 + timeout_seconds: 硬超时时间,None 表示不限制 返回: tuple[bool, str, str]: (是否成功, 标准输出, 标准错误) @@ -104,7 +110,18 @@ async def run_git_command( stderr_lines.append(text) logger.debug(text, LOG_COMMAND) - stdout_bytes, _ = await asyncio.gather(_collect_stdout(process), _read_stderr()) + output_tasks = asyncio.gather(_collect_stdout(process), _read_stderr()) + try: + if timeout_seconds is None: + stdout_bytes, _ = await output_tasks + else: + stdout_bytes, _ = await asyncio.wait_for( + output_tasks, + timeout=timeout_seconds, + ) + except TimeoutError: + await _kill_process_tree(process) + return False, "", f"命令执行超时({timeout_seconds:g} 秒)" await process.wait() stdout = (stdout_bytes or b"").decode("utf-8").strip() @@ -116,6 +133,24 @@ async def run_git_command( return False, "", str(e) +async def _kill_process_tree(process: asyncio.subprocess.Process) -> None: + """终止超时的 Git 进程及其派生的认证、网络辅助进程。""" + with contextlib.suppress(psutil.Error): + parent = psutil.Process(process.pid) + children = parent.children(recursive=True) + for child in reversed(children): + with contextlib.suppress(psutil.Error): + child.kill() + with contextlib.suppress(psutil.Error): + parent.kill() + + if process.returncode is None: + with contextlib.suppress(ProcessLookupError): + process.kill() + with contextlib.suppress(TimeoutError): + await asyncio.wait_for(process.wait(), timeout=5) + + async def _collect_stdout(process: asyncio.subprocess.Process) -> bytes: """收集子进程的全部 stdout 输出。""" assert process.stdout is not None @@ -238,12 +273,40 @@ async def sparse_checkout_clone( if not success: raise RuntimeError(f"配置稀疏路径失败: {err or out}") - # 强制拉取并同步到远端 - success, out, err = await run_git_command( - ["fetch", "--force", "--depth", "1", "origin", branch], temp_path - ) - if not success: - raise RuntimeError(f"fetch 失败: {err or out}") + # 强制拉取并同步到远端。弱网或远端连接卡住时有限重试, + # 防止一次下载永久占住插件商店任务。 + fetch_error = "" + for attempt in range(3): + success, out, err = await run_git_command( + [ + "-c", + "http.lowSpeedLimit=1024", + "-c", + "http.lowSpeedTime=30", + "fetch", + "--force", + "--depth", + "1", + "origin", + branch, + ], + temp_path, + timeout_seconds=60, + ) + if success: + break + fetch_error = err or out + if attempt < 2: + for lock_name in ("shallow.lock", "index.lock"): + with contextlib.suppress(FileNotFoundError, PermissionError): + (temp_path / ".git" / lock_name).unlink() + logger.warning( + f"git fetch 失败,将重试({attempt + 1}/3): {fetch_error}", + LOG_COMMAND, + ) + await asyncio.sleep(2) + else: + raise RuntimeError(f"fetch 失败: {fetch_error}") # 使用远端强制更新本地分支并覆盖工作区 success, out, err = await run_git_command(