diff --git a/pipeline_service/work_env.py b/pipeline_service/work_env.py index 4be5c56..380bfb3 100644 --- a/pipeline_service/work_env.py +++ b/pipeline_service/work_env.py @@ -14,6 +14,7 @@ pipeline_service/work_env.py - 工作环境管理(沙箱的本地/远程切换 """ import os +import re import shlex import shutil import asyncio @@ -79,6 +80,71 @@ def _ssh_common_args(env: dict) -> list: return args +# ── 远程配置安全校验 ─────────────────────────────────────── + +_HOST_RE = re.compile(r"^[a-zA-Z0-9]([a-zA-Z0-9\-\.]*[a-zA-Z0-9])?$") +_USER_RE = re.compile(r"^[a-zA-Z0-9._-]+$") + +# 禁止作为 remote_dir 的系统目录(防 rsync --delete 误删系统文件) +_FORBIDDEN_DIRS = { + "/", "/bin", "/sbin", "/usr", "/etc", "/lib", "/lib64", "/var", + "/root", "/boot", "/dev", "/proc", "/sys", "/tmp", "/opt", "/run", + "/srv", "/mnt", "/media", "/home", +} + + +def _validate_remote(remote_config: dict) -> str: + """校验远程配置,返回错误消息(None 表示通过)。 + + 防注入(host/user/port)与防破坏(remote_dir 不能是系统目录、 + remote_key_path 不能读任意私钥文件)。 + """ + host = (remote_config.get("remote_host") or "").strip() + if not host: + return "缺少 remote_host" + # 允许主机名/IPv4;IPv6 用 [] 包裹后取内部 + _host = host[1:-1] if host.startswith("[") and host.endswith("]") else host + if not _HOST_RE.match(_host): + return "remote_host 格式非法" + + user = (remote_config.get("remote_user") or "").strip() + if not user: + return "缺少 remote_user" + if not _USER_RE.match(user): + return "remote_user 格式非法" + + try: + port = int(remote_config.get("remote_port", 22) or 22) + if not (1 <= port <= 65535): + return "remote_port 必须在 1-65535" + except (ValueError, TypeError): + return "remote_port 必须是数字" + + remote_dir = (remote_config.get("remote_dir") or "").strip() + if not remote_dir: + return "缺少 remote_dir" + if not remote_dir.startswith("/"): + return "remote_dir 必须是绝对路径" + _dir = remote_dir.rstrip("/") + if _dir in _FORBIDDEN_DIRS: + return "remote_dir 不能是系统目录 " + _dir + for d in _FORBIDDEN_DIRS: + if d != "/" and _dir.startswith(d + "/"): + return "remote_dir 不能位于系统目录 " + d + " 下" + + key = (remote_config.get("remote_key_path") or "").strip() + if key: + ssh_dir = os.path.expanduser("~/.ssh") + if key.startswith("/"): + # 绝对路径必须位于 ~/.ssh/ 下,防读任意私钥 + if not key.startswith(ssh_dir.rstrip("/") + "/"): + return "remote_key_path 必须位于 ~/.ssh/ 目录下" + elif "/" in key or key.startswith(".."): + return "remote_key_path 必须是 ~/.ssh/ 下的文件名" + + return None + + # ═══════════════════════════════════════════════════════════ # 目录迁移 # ═══════════════════════════════════════════════════════════ @@ -97,6 +163,18 @@ async def migrate_work_dir(account_name: str, from_mode: str, to_mode: str, env: if not RSYNC: return {"ok": False, "error": "rsync 不可用,无法迁移"} + # remote_dir 安全校验(纵深防御:DB 脏数据/历史记录也拦一道) + if from_mode == "remote" or to_mode == "remote": + remote_dir = (env.get("remote_dir") or "").strip() + if not remote_dir: + return {"ok": False, "error": "缺少远程目录 remote_dir"} + _d = remote_dir.rstrip("/") + if _d in _FORBIDDEN_DIRS: + return {"ok": False, "error": "远程目录是系统目录,拒绝迁移: " + _d} + for d in _FORBIDDEN_DIRS: + if d != "/" and _d.startswith(d + "/"): + return {"ok": False, "error": "远程目录位于系统目录下,拒绝迁移: " + d} + # 构造 rsync 命令 try: if from_mode == "local" and to_mode == "remote": @@ -151,12 +229,9 @@ async def set_work_env(sor, owner_type: str, owner_id: str, mode: str, remote_config = remote_config or {} if mode == "remote": - if not remote_config.get("remote_host"): - return {"ok": False, "error": "远程模式需要 remote_host"} - if not remote_config.get("remote_user"): - return {"ok": False, "error": "远程模式需要 remote_user"} - if not remote_config.get("remote_dir"): - return {"ok": False, "error": "远程模式需要 remote_dir"} + err = _validate_remote(remote_config) + if err: + return {"ok": False, "error": err} # 查旧环境(完整记录,迁移时需要用旧的远程配置) old_recs = await sor.sqlExe( @@ -177,6 +252,20 @@ async def set_work_env(sor, owner_type: str, owner_id: str, mode: str, "remote_key_path": remote_config.get("remote_key_path", ""), "remote_dir": remote_config.get("remote_dir", ""), } + # 先迁移(切换模式时),成功后再更新 DB(保证状态一致,迁移失败不回写 mode) + migrated = {"migrated": False} + if account_name and old_mode != mode: + # remote → local 用旧的远程配置;local → remote 用新的远程配置 + if old_mode == "remote" and mode == "local": + migrate_env = old_env or _row_to_env_from_dict(data) + else: + migrate_env = _row_to_env_from_dict(data) + migrated = await migrate_work_dir(account_name, old_mode, mode, migrate_env) + if not migrated.get("ok"): + return {"ok": False, "error": "迁移失败,未切换模式: " + migrated.get("error", ""), + "migration": migrated} + + # upsert(迁移成功后才写 DB) if old_recs: await sor.sqlExe( "UPDATE sd_work_envs SET mode=${mode}$, remote_host=${h}$, remote_port=${p}$, " @@ -189,16 +278,6 @@ async def set_work_env(sor, owner_type: str, owner_id: str, mode: str, data["id"] = getID() await sor.C("sd_work_envs", data) - # 触发迁移(模式切换 + 提供了 account_name 时) - migrated = {"migrated": False} - if account_name and old_mode != mode: - # remote → local 用旧的远程配置;local → remote 用新的远程配置 - if old_mode == "remote" and mode == "local": - migrate_env = old_env or _row_to_env_from_dict(data) - else: - migrate_env = _row_to_env_from_dict(data) - migrated = await migrate_work_dir(account_name, old_mode, mode, migrate_env) - return {"ok": True, "mode": mode, "old_mode": old_mode, "migration": migrated} @@ -291,6 +370,20 @@ async def run_in_work_env(sor, account_name: str, user_id: str, org_id: str, async def ensure_remote_bwrap(env: dict) -> dict: """确保远程主机已安装 bwrap(零 root:apt download + dpkg -x)。""" + # 校验(host/user/port/key,无需 remote_dir) + host = (env.get("remote_host") or "").strip() + _host = host[1:-1] if host.startswith("[") and host.endswith("]") else host + if not host or not _HOST_RE.match(_host): + return {"ok": False, "error": "remote_host 格式非法"} + user = (env.get("remote_user") or "").strip() + if not user or not _USER_RE.match(user): + return {"ok": False, "error": "remote_user 格式非法"} + key = (env.get("remote_key_path") or "").strip() + if key and key.startswith("/"): + ssh_dir = os.path.expanduser("~/.ssh") + if not key.startswith(ssh_dir.rstrip("/") + "/"): + return {"ok": False, "error": "remote_key_path 必须位于 ~/.ssh/ 目录下"} + ssh_args = _ssh_common_args(env) target = _ssh_target(env)