Files
zhenxun_bot/zhenxun/utils/repo_utils/utils.py
T
ManyManyTomatoandATTomatoo e5fa0f0335 fix(plugin-store): 优化下载逻辑及回退 (#2153)
* feat: 添加插件下载源管理与二进制文件处理逻辑

* fix(test): 修复强制阿里云测试中的参数断言,确保 repo_type 正确

---------

Co-authored-by: ATTomatoo <1126160939@qq.com>
2026-07-30 20:48:44 +08:00

415 lines
13 KiB
Python
Raw Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
"""
仓库管理工具的工具函数
"""
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
from .exceptions import GitUnavailableError
async def check_git() -> bool:
"""
检查环境变量中是否存在 git
返回:
bool: 是否存在git命令
"""
try:
process = await asyncio.create_subprocess_exec(
"git",
"--version",
stdout=asyncio.subprocess.PIPE,
stderr=asyncio.subprocess.PIPE,
)
stdout, _ = await process.communicate()
return bool(stdout)
except Exception as e:
logger.error("检查git命令失败", LOG_COMMAND, e=e)
return False
async def clean_git(cwd: Path):
"""
清理git仓库
参数:
cwd: 工作目录
"""
await run_git_command("reset --hard", cwd)
await run_git_command("clean -xdf", cwd)
async def run_git_command(
command: str | list[str],
cwd: Path | None = None,
timeout_seconds: float | None = None,
) -> tuple[bool, str, str]:
"""
运行git命令,实时输出 stderr 进度信息(如 git clone --progress)。
参数:
command: 命令字符串或参数列表
cwd: 工作目录
timeout_seconds: 硬超时时间,None 表示不限制
返回:
tuple[bool, str, str]: (是否成功, 标准输出, 标准错误)
"""
try:
args = command.split() if isinstance(command, str) else list(command)
process = await asyncio.create_subprocess_exec(
"git",
*args,
stdout=asyncio.subprocess.PIPE,
stderr=asyncio.subprocess.PIPE,
cwd=cwd,
)
stderr_lines: list[str] = []
async def _read_stderr():
assert process.stderr is not None
buf = b""
while True:
chunk = await process.stderr.read(256)
if not chunk:
if buf:
text = buf.decode("utf-8", errors="replace").strip()
if text:
stderr_lines.append(text)
logger.debug(text, LOG_COMMAND)
break
buf += chunk
while b"\n" in buf or b"\r" in buf:
idx_r = buf.find(b"\r")
idx_n = buf.find(b"\n")
if idx_r == -1:
idx = idx_n
elif idx_n == -1:
idx = idx_r
else:
idx = min(idx_r, idx_n)
line_bytes = buf[:idx]
if buf[idx : idx + 2] == b"\r\n":
buf = buf[idx + 2 :]
else:
buf = buf[idx + 1 :]
text = line_bytes.decode("utf-8", errors="replace").strip()
if text:
stderr_lines.append(text)
logger.debug(text, LOG_COMMAND)
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()
stderr = "\n".join(stderr_lines)
return process.returncode == 0, stdout, stderr
except Exception as e:
logger.error(f"运行git命令失败: {command}, 错误: {e}")
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
chunks: list[bytes] = []
while True:
chunk = await process.stdout.read(4096)
if not chunk:
break
chunks.append(chunk)
return b"".join(chunks)
def glob_to_regex(pattern: str) -> str:
"""
将glob模式转换为正则表达式
参数:
pattern: glob模式,如 "*.py"
返回:
str: 正则表达式
"""
# 转义特殊字符
regex = re.escape(pattern)
# 替换glob通配符
regex = regex.replace(r"\*\*", ".*") # ** -> .*
regex = regex.replace(r"\*", "[^/]*") # * -> [^/]*
regex = regex.replace(r"\?", "[^/]") # ? -> [^/]
# 添加开始和结束标记
regex = f"^{regex}$"
return regex
def filter_files(
files: list[str],
include_patterns: list[str] | None = None,
exclude_patterns: list[str] | None = None,
) -> list[str]:
"""
过滤文件列表
参数:
files: 文件列表
include_patterns: 包含的文件模式列表,如 ["*.py", "docs/*.md"]
exclude_patterns: 排除的文件模式列表,如 ["__pycache__/*", "*.pyc"]
返回:
list[str]: 过滤后的文件列表
"""
result = files.copy()
# 应用包含模式
if include_patterns:
included = []
for pattern in include_patterns:
regex_pattern = glob_to_regex(pattern)
included.extend(file for file in result if re.match(regex_pattern, file))
result = included
# 应用排除模式
if exclude_patterns:
for pattern in exclude_patterns:
regex_pattern = glob_to_regex(pattern)
result = [file for file in result if not re.match(regex_pattern, file)]
return result
async def sparse_checkout_clone(
repo_url: str,
branch: str,
sparse_path: str | list[str],
target_dir: Path,
) -> list[str]:
"""
使用 git 稀疏检出克隆指定路径到目标目录(在临时目录中操作)。
关键保障:
- 在临时目录中执行所有 git 操作,避免影响 target_dir 中的现有内容
- 只操作 target_dir/sparse_path 路径,不影响 target_dir 其他内容
返回:
list[str]: 成功检出的路径
"""
target_dir.mkdir(parents=True, exist_ok=True)
sparse_paths = [sparse_path] if isinstance(sparse_path, str) else list(sparse_path)
if not await check_git():
raise GitUnavailableError()
# 在临时目录中进行 git 操作
with tempfile.TemporaryDirectory() as temp_dir:
temp_path = Path(temp_dir)
# 初始化临时目录为 git 仓库
success, out, err = await run_git_command("init", temp_path)
if not success:
raise RuntimeError(f"git init 失败: {err or out}")
success, out, err = await run_git_command(
["remote", "add", "origin", repo_url], temp_path
)
if not success:
raise RuntimeError(f"添加远程失败: {err or out}")
# 启用稀疏检出(使用 --no-cone 模式以获得更精确的控制)
await run_git_command("config core.sparseCheckout true", temp_path)
await run_git_command("sparse-checkout init --no-cone", temp_path)
# 设置需要检出的路径(每次都覆盖配置)
if not sparse_paths or any(not path for path in sparse_paths):
raise RuntimeError("sparse-checkout 路径不能为空")
# 使用 --no-cone 模式,直接指定要检出的具体路径
success, out, err = await run_git_command(
["sparse-checkout", "set", "--", *sparse_paths], temp_path
)
if not success:
raise RuntimeError(f"配置稀疏路径失败: {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(
["checkout", "-B", branch, f"origin/{branch}"], temp_path
)
if not success:
# 回退方案
success2, out2, err2 = await run_git_command(
["checkout", branch], temp_path
)
if not success2:
raise RuntimeError(f"checkout 失败: {(err or out) or (err2 or out2)}")
# 强制对齐工作区
await run_git_command(["reset", "--hard", f"origin/{branch}"], temp_path)
await run_git_command("clean -xdf", temp_path)
# 将成功检出的文件复制到 staging,缺失路径由调用方决定是否忽略
downloaded_paths = []
for path in sparse_paths:
source_path = temp_path / path
if not source_path.exists():
continue
target_path = target_dir / path
target_path.parent.mkdir(parents=True, exist_ok=True)
if target_path.exists():
if target_path.is_dir():
shutil.rmtree(target_path)
else:
target_path.unlink()
if source_path.is_dir():
shutil.copytree(source_path, target_path)
else:
shutil.copy2(source_path, target_path)
downloaded_paths.append(path)
return downloaded_paths
def prepare_aliyun_url(repo_url: str, group_name: str | None = None) -> str:
"""解析阿里云CodeUp的仓库URL
参数:
repo_url: 仓库URL
group_name: 分组名称,如果为None则使用默认组织名称
返回:
str: 解析后的仓库URL
"""
config = RepoConfig.get_instance()
repo_name = repo_url.split("/tree/")[0].split("/")[-1].replace(".git", "")
# 使用指定的分组名或默认组织名称
group = group_name or config.aliyun_codeup.organization_name
# 构建仓库URL
# 阿里云CodeUp的仓库URL格式通常为:
# https://codeup.aliyun.com/{organization_id}/{group_name}/{repo_name}.git
url = f"https://codeup.aliyun.com/{config.aliyun_codeup.organization_id}/{group}/{repo_name}.git"
# 添加访问令牌 - 使用base64解码后的令牌
if config.aliyun_codeup.rdc_access_token_encrypted:
try:
# 解码RDC访问令牌
token = base64.b64decode(
config.aliyun_codeup.rdc_access_token_encrypted.encode()
).decode()
# 阿里云CodeUp使用oauth2:token的格式进行身份验证
url = url.replace("https://", f"https://oauth2:{token}@")
logger.debug(f"使用RDC令牌构建阿里云URL: {url.split('@')[0]}@***")
except Exception as e:
logger.error(f"解码RDC令牌失败: {e}")
return url
async def get_aliyun_group_for_repo(repo_name: str) -> str | None:
"""获取仓库所属的阿里云分组名
参数:
repo_name: 仓库名称
返回:
str | None: 分组名称,如果在核心映射中则返回None(使用默认组织名)
"""
from zhenxun.utils.github_utils.const import (
ALIYUN_EXTERNAL_PLUGIN_GROUPS,
ALIYUN_REPO_MAPPING,
)
from zhenxun.utils.github_utils.models import AliyunFileInfo
# 如果在核心映射中,使用默认组织名
if repo_name in ALIYUN_REPO_MAPPING:
return None
# 尝试从外部插件分组中查找
for group_path in ALIYUN_EXTERNAL_PLUGIN_GROUPS:
try:
repos = await AliyunFileInfo.list_group_repositories(group_path)
for repo in repos:
if repo.get("name") == repo_name:
return group_path
except Exception:
continue
return None