""" 仓库管理工具的工具函数 """ 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