Merge pull request #11 from jbwfu/fix/repo-sync

refactor: 优化 sync.py 脚本逻辑,支持分支自动检测与类型检查
This commit is contained in:
engigu
2026-02-07 11:08:48 +08:00
committed by GitHub
+209 -194
View File
@@ -5,30 +5,103 @@ import argparse
import os import os
import subprocess import subprocess
import sys import sys
import shutil
import urllib.request
import urllib.error import urllib.error
import urllib.request
from typing import TYPE_CHECKING, Protocol, cast
# 仅在类型检查时导入,避免运行时依赖
if TYPE_CHECKING:
from http.client import HTTPResponse
def run(cmd, env=None, cwd=None): class SyncArgs(Protocol):
"""执行命令并打印输出""" """用于类型检查的参数协议,映射 argparse 的解析结果"""
print(">>", " ".join(cmd)) source_type: str
result = subprocess.run( source_url: str
cmd, target_path: str
cwd=cwd, branch: str | None
env=env, path: str | None
stdout=sys.stdout, single_file: bool
stderr=sys.stderr, proxy: str | None
) proxy_url: str | None
auth_token: str | None
http_proxy: str | None
def run(
cmd: list[str],
env: dict[str, str] | None = None,
cwd: str | None = None,
capture_output: bool = False
) -> str | None:
"""
执行系统命令。
Args:
cmd: 命令列表
env: 环境变量字典
cwd: 当前工作目录
capture_output: 是否捕获输出。
"""
if not capture_output:
print(">>", " ".join(cmd))
if capture_output:
# 捕获模式:不打印到屏幕,返回 stdout
result = subprocess.run(
cmd,
cwd=cwd,
env=env,
stdout=subprocess.PIPE,
stderr=subprocess.PIPE,
timeout=30,
encoding="utf-8",
errors="ignore"
)
else:
# 直通模式:直接打印到屏幕
result = subprocess.run(
cmd,
cwd=cwd,
env=env,
stdout=sys.stdout,
stderr=sys.stderr,
)
if result.returncode != 0: if result.returncode != 0:
if capture_output:
return None
sys.exit(result.returncode) sys.exit(result.returncode)
if capture_output:
return str(result.stdout).strip()
def build_proxy_url(url, proxy_type, proxy_url): return str(result)
"""构建代理 URL"""
def get_remote_default_branch(repo_url: str, env: dict[str, str]) -> str:
"""检测远程仓库的默认分支名称(通常是 main 或 master)。"""
print(f"正在检测远程仓库默认分支: {repo_url}")
cmd = ["git", "ls-remote", "--symref", repo_url, "HEAD"]
output = run(cmd, env=env, capture_output=True)
if isinstance(output, str):
for line in output.splitlines():
parts = line.split()
if len(parts) >= 2 and parts[0] == "ref:" and "refs/heads/" in parts[1]:
branch = parts[1].removeprefix("refs/heads/")
print(f"检测到默认分支: {branch}")
return branch
print("无法检测到默认分支,回退使用 'main'")
return "main"
def build_proxy_url(url: str, proxy_type: str | None, proxy_url: str | None) -> str:
"""根据配置构建带有代理前缀的 URL。"""
if not proxy_type or proxy_type == "none": if not proxy_type or proxy_type == "none":
return url return url
proxy_base = "" proxy_base = ""
if proxy_type == "ghproxy": if proxy_type == "ghproxy":
proxy_base = "https://gh-proxy.com/" proxy_base = "https://gh-proxy.com/"
@@ -36,72 +109,37 @@ def build_proxy_url(url, proxy_type, proxy_url):
proxy_base = "https://mirror.ghproxy.com/" proxy_base = "https://mirror.ghproxy.com/"
elif proxy_type == "custom" and proxy_url: elif proxy_type == "custom" and proxy_url:
proxy_base = proxy_url.rstrip("/") + "/" proxy_base = proxy_url.rstrip("/") + "/"
if proxy_base and url.startswith("http"): if proxy_base and url.startswith("http"):
return proxy_base + url return proxy_base + url
return url return url
def sync_git_file(args, repo_url, env): def _download_file(url: str, dest: str, auth_token: str | None) -> None:
"""从 Git 仓库同步单个文件(通过 raw URL 下载)""" """
source_url = args.source_url 内部通用下载函数,处理请求构建、Token 认证和文件写入。
file_path = args.path """
branch = args.branch or "main" print(f"下载地址: {url}")
dest = args.target_path
# 构建 raw 文件 URL
# GitHub: https://github.com/user/repo -> https://raw.githubusercontent.com/user/repo/branch/path
# GitLab: https://gitlab.com/user/repo -> https://gitlab.com/user/repo/-/raw/branch/path
# Gitee: https://gitee.com/user/repo -> https://gitee.com/user/repo/raw/branch/path
raw_url = None
if "github.com" in source_url:
# GitHub
base = source_url.replace("github.com", "raw.githubusercontent.com").rstrip(".git")
raw_url = f"{base}/{branch}/{file_path}"
elif "gitlab.com" in source_url:
# GitLab
base = source_url.rstrip(".git")
raw_url = f"{base}/-/raw/{branch}/{file_path}"
elif "gitee.com" in source_url:
# Gitee
base = source_url.rstrip(".git")
raw_url = f"{base}/raw/{branch}/{file_path}"
else:
# 通用:尝试 GitHub 风格
base = source_url.rstrip(".git")
raw_url = f"{base}/raw/{branch}/{file_path}"
# 应用代理
raw_url = build_proxy_url(raw_url, args.proxy, args.proxy_url)
print(f"下载单文件: {raw_url}")
print(f"目标路径: {dest}") print(f"目标路径: {dest}")
# 确保目标目录存在
parent_dir = os.path.dirname(dest) parent_dir = os.path.dirname(dest)
if parent_dir: if parent_dir:
os.makedirs(parent_dir, exist_ok=True) os.makedirs(parent_dir, exist_ok=True)
# 创建请求 req = urllib.request.Request(url)
req = urllib.request.Request(raw_url) if auth_token:
req.add_header("Authorization", f"token {auth_token}")
# 添加认证 Token
if args.auth_token:
req.add_header("Authorization", f"token {args.auth_token}")
req.add_header("User-Agent", "Mozilla/5.0 (compatible; sync.py)") req.add_header("User-Agent", "Mozilla/5.0 (compatible; sync.py)")
try: try:
with urllib.request.urlopen(req, timeout=300) as response: with cast("HTTPResponse", urllib.request.urlopen(req, timeout=300)) as response:
content = response.read() content: bytes = response.read()
with open(dest, "wb") as f: with open(dest, "wb") as f:
f.write(content) _ = f.write(content)
print(f"文件大小: {len(content)} 字节") print(f"文件大小: {len(content)} 字节")
print("同步完成") print("下载完成")
except urllib.error.HTTPError as e: except urllib.error.HTTPError as e:
print(f"下载失败, HTTP 状态码: {e.code}") print(f"下载失败, HTTP 状态码: {e.code}")
sys.exit(1) sys.exit(1)
@@ -110,8 +148,49 @@ def sync_git_file(args, repo_url, env):
sys.exit(1) sys.exit(1)
def is_raw_file_url(url): def sync_git_file(args: SyncArgs, repo_url: str, env: dict[str, str]) -> None:
"""检测是否是 raw 文件 URL""" """
从 Git 仓库同步单个文件(通过构造 Raw URL 下载)。
"""
source_url = args.source_url
file_path = args.path or ""
dest = args.target_path
# 如果目标是目录,自动拼接文件名
if os.path.isdir(dest) or dest.endswith(os.sep) or (os.altsep and dest.endswith(os.altsep)):
filename = os.path.basename(file_path)
dest = os.path.join(dest, filename)
print(f"检测到目标路径为目录 '{args.target_path}',自动修正为: '{dest}'")
branch = args.branch or get_remote_default_branch(repo_url, env)
# 构建 raw 文件 URL
# GitHub: https://github.com/user/repo -> https://raw.githubusercontent.com/user/repo/branch/path
# GitLab: https://gitlab.com/user/repo -> https://gitlab.com/user/repo/-/raw/branch/path
# Gitee: https://gitee.com/user/repo -> https://gitee.com/user/repo/raw/branch/path
clean_url = args.source_url.rstrip(".git")
raw_url = ""
if "github.com" in source_url:
base = args.source_url.replace("github.com", "raw.githubusercontent.com").rstrip(".git")
raw_url = f"{base}/{branch}/{file_path}"
elif "gitlab.com" in source_url:
raw_url = f"{clean_url}/-/raw/{branch}/{file_path}"
elif "gitee.com" in source_url:
raw_url = f"{clean_url}/raw/{branch}/{file_path}"
else:
# 通用策略:尝试 GitHub 风格
raw_url = f"{clean_url}/raw/{branch}/{file_path}"
raw_url = build_proxy_url(raw_url, args.proxy, args.proxy_url)
# 调用通用下载函数
_download_file(raw_url, dest, args.auth_token)
def is_raw_file_url(url: str) -> bool:
"""判断 URL 是否已经是 Raw 文件链接。"""
raw_patterns = [ raw_patterns = [
"raw.githubusercontent.com", "raw.githubusercontent.com",
"/raw/", "/raw/",
@@ -121,45 +200,38 @@ def is_raw_file_url(url):
return any(pattern in url for pattern in raw_patterns) return any(pattern in url for pattern in raw_patterns)
def get_repo_name(url): def get_repo_name(url: str) -> str:
"""从 Git URL 中提取仓库名""" """从 Git URL 中提取仓库名称。"""
# 去掉末尾的 .git url_stripped = url.rstrip("/").rstrip(".git")
url = url.rstrip("/").rstrip(".git") return os.path.basename(url_stripped)
# 提取最后一部分作为仓库名
return os.path.basename(url)
def sync_git(args): def sync_git(args: SyncArgs) -> None:
"""Git 仓库同步""" """处理 Git 类型的同步逻辑(Clone, Pull 或 Sparse Checkout)。"""
env = os.environ.copy() env = os.environ.copy()
# 如果 source_url 是 raw 文件 URL,自动切换到 URL 下载模式
if is_raw_file_url(args.source_url): if is_raw_file_url(args.source_url):
print("检测到 raw 文件 URL,自动切换到 URL 下载模式") print("检测到 raw 文件 URL,自动切换到 URL 下载模式")
sync_url(args) sync_url(args)
return return
# 设置 HTTP 代理
if args.http_proxy: if args.http_proxy:
env["http_proxy"] = args.http_proxy env["http_proxy"] = args.http_proxy
env["https_proxy"] = args.http_proxy env["https_proxy"] = args.http_proxy
# 构建仓库 URL(带代理)
repo_url = build_proxy_url(args.source_url, args.proxy, args.proxy_url) repo_url = build_proxy_url(args.source_url, args.proxy, args.proxy_url)
# 如果有认证 Token,将其嵌入 URL
if args.auth_token and repo_url.startswith("https://"): if args.auth_token and repo_url.startswith("https://"):
repo_url = repo_url.replace("https://", f"https://{args.auth_token}@") repo_url = repo_url.replace("https://", f"https://{args.auth_token}@")
dest = args.target_path dest = args.target_path
branch = args.branch or "main"
# 如果指定了 path 且是单文件模式使用 raw URL 下载 # 单文件模式使用 Raw URL 下载
if args.path and args.single_file: if args.path and args.single_file:
sync_git_file(args, repo_url, env) sync_git_file(args, repo_url, env)
return return
# 如果目标路径是已存在的目录且不是 git 仓库,自动追加仓库名作为子目录 # 自动追加仓库名逻辑
git_dir = os.path.join(dest, ".git") git_dir = os.path.join(dest, ".git")
if os.path.isdir(dest) and not os.path.exists(git_dir): if os.path.isdir(dest) and not os.path.exists(git_dir):
repo_name = get_repo_name(args.source_url) repo_name = get_repo_name(args.source_url)
@@ -167,143 +239,86 @@ def sync_git(args):
print(f"目标路径自动追加仓库名: {dest}") print(f"目标路径自动追加仓库名: {dest}")
git_dir = os.path.join(dest, ".git") git_dir = os.path.join(dest, ".git")
# 检查目标目录是否已存在 git 仓库 if os.path.exists(git_dir):
is_existing_repo = os.path.exists(git_dir) print("检测到已存在仓库,执行 git pull")
if args.branch:
if is_existing_repo: _ = run(["git", "checkout", args.branch], cwd=dest, env=env)
# 已存在仓库,执行 git pull _ = run(["git", "pull"], cwd=dest, env=env)
print(f"检测到已存在仓库,执行 git pull")
# 先切换分支
if branch:
try:
run(["git", "checkout", branch], cwd=dest, env=env)
except:
pass
run(["git", "pull"], cwd=dest, env=env)
else: else:
# 新仓库,执行 git clone print("执行 git clone")
print(f"执行 git clone")
# 确保父目录存在
parent_dir = os.path.dirname(dest) parent_dir = os.path.dirname(dest)
if parent_dir: if parent_dir:
os.makedirs(parent_dir, exist_ok=True) os.makedirs(parent_dir, exist_ok=True)
# 如果目标目录已存在且不为空,报错提示
if os.path.exists(dest) and os.listdir(dest): if os.path.exists(dest) and os.listdir(dest):
print(f"错误: 目标目录 '{dest}' 已存在且不为空,无法执行 git clone") print(f"错误: 目标目录 '{dest}' 已存在且不为空,无法执行 git clone")
print("提示: 请清空目标目录或指定一个新目录") print("提示: 请清空目标目录或指定一个新目录")
sys.exit(1) sys.exit(1)
# 稀疏 clone(如果指定了 path) clone_cmd = ["git", "clone", "--depth", "1"]
if args.path:
run([
"git", "clone",
"--depth", "1",
"--filter=blob:none",
"--no-checkout",
"-b", branch,
repo_url,
dest
], env=env)
run(["git", "sparse-checkout", "init", "--cone"], cwd=dest, env=env) if args.branch:
run(["git", "sparse-checkout", "set", args.path], cwd=dest, env=env) clone_cmd.extend(["-b", args.branch])
run(["git", "checkout"], cwd=dest, env=env)
# 稀疏检出 (Sparse Checkout)
if args.path:
clone_cmd.extend(["--filter=blob:none", "--no-checkout", repo_url, dest])
_ = run(clone_cmd, env=env)
_ = run(["git", "sparse-checkout", "init", "--cone"], cwd=dest, env=env)
_ = run(["git", "sparse-checkout", "set", args.path], cwd=dest, env=env)
_ = run(["git", "checkout"], cwd=dest, env=env)
else: else:
# 普通 clone # 普通 Clone
run([ clone_cmd.extend([repo_url, dest])
"git", "clone", _ = run(clone_cmd, env=env)
"--depth", "1",
"-b", branch,
repo_url,
dest
], env=env)
print("同步完成") print("同步完成")
def sync_url(args): def sync_url(args: SyncArgs) -> None:
"""URL 文件下载""" """处理普通 URL 文件下载逻辑。"""
# 构建下载 URL(带代理)
download_url = build_proxy_url(args.source_url, args.proxy, args.proxy_url) download_url = build_proxy_url(args.source_url, args.proxy, args.proxy_url)
print(f"下载地址: {download_url}") print(f"下载地址: {download_url}")
dest = args.target_path dest = args.target_path
# 如果目标路径是目录或以 / 结尾,从 URL 中提取文件名
if os.path.isdir(dest) or dest.endswith("/"): if os.path.isdir(dest) or dest.endswith("/"):
# 从 URL 中提取文件名 url_path = args.source_url.split("?")[0]
url_path = args.source_url.split("?")[0] # 去掉查询参数 filename = os.path.basename(url_path) or "downloaded_file"
filename = os.path.basename(url_path)
if not filename:
filename = "downloaded_file"
dest = os.path.join(dest, filename) dest = os.path.join(dest, filename)
print(f"目标文件: {dest}") print(f"目标文件: {dest}")
# 确保目标目录存在
parent_dir = os.path.dirname(dest)
if parent_dir:
os.makedirs(parent_dir, exist_ok=True)
# 创建请求 # 调用通用下载函数
req = urllib.request.Request(download_url) _download_file(download_url, dest, args.auth_token)
# 添加认证 Token
if args.auth_token:
req.add_header("Authorization", f"token {args.auth_token}")
# 添加 User-Agent
req.add_header("User-Agent", "Mozilla/5.0 (compatible; sync.py)")
try:
with urllib.request.urlopen(req, timeout=300) as response:
content = response.read()
with open(dest, "wb") as f:
f.write(content)
print(f"目标路径: {dest}")
print(f"文件大小: {len(content)} 字节")
print("同步完成")
except urllib.error.HTTPError as e:
print(f"下载失败, HTTP 状态码: {e.code}")
sys.exit(1)
except urllib.error.URLError as e:
print(f"下载失败: {e.reason}")
sys.exit(1)
def main(): def main() -> None:
parser = argparse.ArgumentParser(description="仓库/文件同步工具") parser = argparse.ArgumentParser(description="仓库/文件同步工具")
parser.add_argument("--source-type", choices=["git", "url"], default="git", _ = parser.add_argument("--source-type", choices=["git", "url"], default="git",
help="源类型: git(Git仓库) 或 url(URL下载)") help="源类型: git(Git仓库) 或 url(URL下载)")
parser.add_argument("--source-url", required=True, _ = parser.add_argument("--source-url", required=True,
help="源地址(Git仓库URL或文件URL") help="源地址(Git仓库URL或文件URL")
parser.add_argument("--target-path", required=True, _ = parser.add_argument("--target-path", required=True,
help="目标路径") help="目标路径")
parser.add_argument("--branch", default="main", _ = parser.add_argument("--branch",
help="Git 分支名(仅 git 类型有效)") help="Git 分支名(仅 git 类型有效)")
parser.add_argument("--path", _ = parser.add_argument("--path",
help="仅拉取指定文件或目录(仅 git 类型有效)") help="仅拉取指定文件或目录(仅 git 类型有效)")
parser.add_argument("--single-file", action="store_true", _ = parser.add_argument("--single-file", action="store_true",
help="单文件模式,直接下载指定文件而非 sparse-checkout(需配合 --path 使用)") help="单文件模式,直接下载指定文件而非 sparse-checkout(需配合 --path 使用)")
parser.add_argument("--proxy", choices=["none", "ghproxy", "mirror", "custom"], default="none", _ = parser.add_argument("--proxy", choices=["none", "ghproxy", "mirror", "custom"], default="none",
help="代理类型") help="代理类型")
parser.add_argument("--proxy-url", _ = parser.add_argument("--proxy-url",
help="自定义代理地址(仅 proxy=custom 时有效)") help="自定义代理地址(仅 proxy=custom 时有效)")
parser.add_argument("--auth-token", _ = parser.add_argument("--auth-token",
help="认证 Token(用于私有仓库)") help="认证 Token(用于私有仓库)")
parser.add_argument("--http-proxy", _ = parser.add_argument("--http-proxy",
help="HTTP 代理(如 http://127.0.0.1:7890") help="HTTP 代理(如 http://127.0.0.1:7890")
args = parser.parse_args() # 使用 cast 将 Namespace 转换为 SyncArgs 协议,满足静态类型检查
args = cast(SyncArgs, cast(object, parser.parse_args()))
# 打印原始命令行参数
print("参数:", " ".join(sys.argv[1:])) print("参数:", " ".join(sys.argv[1:]))
if args.source_type == "git": if args.source_type == "git":