Skip to content

Instantly share code, notes, and snippets.

@alanzchen
Last active July 24, 2026 15:32
Show Gist options
  • Select an option

  • Save alanzchen/451de5399e543b7d5429831c9d59ca94 to your computer and use it in GitHub Desktop.

Select an option

Save alanzchen/451de5399e543b7d5429831c9d59ca94 to your computer and use it in GitHub Desktop.
Cloudflare-only Colab SSH worker init script
#!/usr/bin/env python3
"""Cloudflare-only Colab SSH worker bootstrap.
Designed to be safe for a public GitHub gist. All deployment-specific values are
read from environment variables:
Required:
COLAB_WORKER_PUBLIC_KEY
SSH public key to authorize for root login in the Colab runtime.
Optional:
COLAB_WORKER_HOSTNAME_PREFIX
Prefix for the printed runtime hostname. Default: colab-worker.
COLAB_WORKER_SSH_PORT
Local sshd port inside Colab. Default: 2222.
COLAB_WORKER_MOUNT_DRIVE
Whether to mount Google Drive. Default: 0.
COLAB_WORKER_DRIVE_PROJECT_DIRS
Colon-separated candidate Drive project directories inside Colab.
COLAB_WORKER_SNAPSHOT_REL_PATH
Relative path, under the first existing Drive project dir, to a .tgz
workspace snapshot to unpack.
COLAB_WORKER_WORKSPACE
Workspace extraction target. Default: /content/workspace.
COLAB_WORKER_CLOUDFLARED_URL
cloudflared Linux amd64 download URL.
"""
from __future__ import annotations
import os
import pathlib
import platform
import re
import shlex
import subprocess
import time
def getenv_bool(name: str, default: bool) -> bool:
value = os.environ.get(name)
if value is None:
return default
return value.strip().lower() not in {"0", "false", "no", "off", ""}
def run(cmd: str, check: bool = True) -> subprocess.CompletedProcess[str]:
print(f"\n$ {cmd}", flush=True)
return subprocess.run(cmd, shell=True, check=check, text=True)
def prepend_search_paths(current: str, required: list[str]) -> str:
"""Prepend required paths without duplicating existing entries."""
output: list[str] = []
for item in [*required, *current.split(os.pathsep)]:
item = item.strip()
if item and item not in output:
output.append(item)
return os.pathsep.join(output)
def sshd_setenv_option() -> str:
"""Build an explicit environment inherited by every SSH command.
Colab can attach a GPU after this launcher starts. Keeping the standard
CUDA and NVIDIA paths in sshd's session environment makes that late-bound
GPU visible without restarting the tunnel.
"""
ssh_path = prepend_search_paths(
os.environ.get(
"PATH",
"/usr/local/sbin:/usr/local/bin:/usr/sbin:/usr/bin:/sbin:/bin",
),
["/usr/local/cuda/bin"],
)
ssh_library_path = prepend_search_paths(
os.environ.get("LD_LIBRARY_PATH", ""),
["/usr/lib64-nvidia", "/usr/local/cuda/lib64"],
)
cuda_home = os.environ.get("CUDA_HOME", "/usr/local/cuda").strip()
values = {
"PATH": ssh_path,
"LD_LIBRARY_PATH": ssh_library_path,
"CUDA_HOME": cuda_home,
}
for name, value in values.items():
if not value or any(char.isspace() for char in value):
raise SystemExit(f"Unsafe or empty SSH environment value for {name}")
assignments = " ".join(f"{name}={value}" for name, value in values.items())
return f"SetEnv={assignments}"
public_key = os.environ.get("COLAB_WORKER_PUBLIC_KEY", "").strip()
if not public_key:
raise SystemExit("Missing required env var: COLAB_WORKER_PUBLIC_KEY")
hostname_prefix = os.environ.get("COLAB_WORKER_HOSTNAME_PREFIX", "colab-worker")
ssh_port = int(os.environ.get("COLAB_WORKER_SSH_PORT", "2222"))
mount_drive = getenv_bool("COLAB_WORKER_MOUNT_DRIVE", False)
workspace = pathlib.Path(
os.environ.get("COLAB_WORKER_WORKSPACE", "/content/workspace")
)
snapshot_rel_path = os.environ.get("COLAB_WORKER_SNAPSHOT_REL_PATH", "").strip()
cloudflared_url = os.environ.get(
"COLAB_WORKER_CLOUDFLARED_URL",
"https://github.com/cloudflare/cloudflared/releases/latest/download/cloudflared-linux-amd64",
)
hostname = f"{hostname_prefix}-{int(time.time())}"
run("apt-get update -qq")
run("DEBIAN_FRONTEND=noninteractive apt-get install -y -qq openssh-server curl ca-certificates rsync")
ssh_dir = pathlib.Path("/root/.ssh")
ssh_dir.mkdir(mode=0o700, exist_ok=True)
authorized_keys = ssh_dir / "authorized_keys"
existing_keys = authorized_keys.read_text() if authorized_keys.exists() else ""
if public_key not in existing_keys:
authorized_keys.write_text(existing_keys.rstrip() + "\n" + public_key + "\n")
authorized_keys.chmod(0o600)
run("ssh-keygen -A")
run("mkdir -p /run/sshd")
run("pkill -x sshd || true", check=False)
ssh_session_environment = sshd_setenv_option()
run(
f"/usr/sbin/sshd -D -p {ssh_port} "
"-o PermitRootLogin=yes "
"-o PasswordAuthentication=no "
"-o PubkeyAuthentication=yes "
"-o AuthorizedKeysFile=/root/.ssh/authorized_keys "
"-o UsePAM=no "
f"-o {shlex.quote(ssh_session_environment)} "
"> /tmp/colab-sshd.log 2>&1 &"
)
drive_project = ""
if mount_drive:
try:
from google.colab import drive
drive.mount("/content/drive")
except Exception as exc: # pragma: no cover - only runs inside Colab.
print(f"WARN: Google Drive mount failed: {exc}", flush=True)
drive_dirs = [
item.strip()
for item in os.environ.get("COLAB_WORKER_DRIVE_PROJECT_DIRS", "").split(":")
if item.strip()
]
for candidate in drive_dirs:
if pathlib.Path(candidate).exists():
drive_project = candidate
break
if drive_project and snapshot_rel_path:
snapshot = pathlib.Path(drive_project) / snapshot_rel_path
if snapshot.exists():
run(f"rm -rf {shlex.quote(str(workspace))} && mkdir -p {shlex.quote(str(workspace))}")
run(f"tar -xzf {shlex.quote(str(snapshot))} -C {shlex.quote(str(workspace))}")
else:
print(f"WARN: snapshot not found at {snapshot}", flush=True)
elif snapshot_rel_path:
print("WARN: snapshot path set but no COLAB_WORKER_DRIVE_PROJECT_DIRS candidate exists.", flush=True)
if platform.machine() not in ("x86_64", "amd64"):
print(f"WARN: unexpected architecture for bundled cloudflared URL: {platform.machine()}", flush=True)
if not pathlib.Path("/usr/local/bin/cloudflared").exists():
run(f"curl -fsSL -o /usr/local/bin/cloudflared {cloudflared_url!r}")
run("chmod +x /usr/local/bin/cloudflared")
run("pkill -x cloudflared || true", check=False)
run("rm -f /tmp/cloudflared-ssh.log")
run(
"nohup /usr/local/bin/cloudflared tunnel --no-autoupdate "
f"--url ssh://127.0.0.1:{ssh_port} > /tmp/cloudflared-ssh.log 2>&1 &"
)
cf_host = ""
for _ in range(90):
log_path = pathlib.Path("/tmp/cloudflared-ssh.log")
log = log_path.read_text(errors="ignore") if log_path.exists() else ""
matches = re.findall(r"https://[A-Za-z0-9.-]+\.trycloudflare\.com", log)
if matches:
cf_host = matches[-1]
break
time.sleep(1)
if not cf_host:
print("WARN: Cloudflare hostname not found. Log follows:", flush=True)
print(pathlib.Path("/tmp/cloudflared-ssh.log").read_text(errors="ignore"), flush=True)
print("\nREADY")
print(f"HOSTNAME={hostname}")
if cf_host:
print(f"CLOUDFLARE_HOST={cf_host}")
if drive_project:
print(f"DRIVE_PROJECT_DIR={drive_project}")
print(f"WORKSPACE={workspace}")
Sign up for free to join this conversation on GitHub. Already have an account? Sign in to comment