mirror of
https://github.com/zhenxun-org/zhenxun_bot.git
synced 2026-09-28 16:20:56 +08:00
* feat: 添加插件下载源管理与二进制文件处理逻辑 * fix(test): 修复强制阿里云测试中的参数断言,确保 repo_type 正确 --------- Co-authored-by: ATTomatoo <1126160939@qq.com>
415 lines
13 KiB
Python
415 lines
13 KiB
Python
"""
|
||
仓库管理工具的工具函数
|
||
"""
|
||
|
||
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
|