LiveHouse-TS / scripts /community_model_wizard.py
ziyuzhou02's picture
Deploy GitHub main 3feb6cda1511
e317359 verified
Raw History Blame Contribute Delete
9.72 kB
#!/usr/bin/env python3
"""Initialize and publish a TS-Live community endpoint with fewer manual steps."""
from __future__ import annotations
import argparse
import json
import shutil
import subprocess
import sys
from datetime import datetime, timezone
from pathlib import Path
from typing import Callable, Sequence
import yaml
REPO_ROOT = Path(__file__).resolve().parents[1]
if str(REPO_ROOT) not in sys.path:
sys.path.insert(0, str(REPO_ROOT))
from scripts.community_endpoint_protocol import validate_endpoint
from scripts.generate_community_model_metadata import (
build_metadata,
validate_submission_identity,
)
TEMPLATE_DIR = REPO_ROOT / "examples" / "community_endpoint"
TEMPLATE_FILES = (
"app.py",
"Dockerfile",
"requirements.txt",
"README.md",
".dockerignore",
"compose.yaml",
)
def initialize_service(output_dir: Path, *, force: bool = False) -> list[Path]:
"""Copy the portable endpoint template without clobbering user work."""
output_dir = output_dir.resolve()
conflicts = [output_dir / name for name in TEMPLATE_FILES if (output_dir / name).exists()]
if conflicts and not force:
names = ", ".join(path.name for path in conflicts)
raise RuntimeError(
f"refusing to overwrite existing files in {output_dir}: {names}; "
"pass --force only if this is intentional"
)
output_dir.mkdir(parents=True, exist_ok=True)
copied = []
for name in TEMPLATE_FILES:
destination = output_dir / name
shutil.copy2(TEMPLATE_DIR / name, destination)
copied.append(destination)
return copied
def _run_command(
command: Sequence[str],
*,
runner: Callable[..., subprocess.CompletedProcess[str]] = subprocess.run,
) -> subprocess.CompletedProcess[str]:
try:
result = runner(
list(command),
check=False,
capture_output=True,
text=True,
)
except FileNotFoundError as exc:
raise RuntimeError(
f"{command[0]!r} was not found; install Tailscale and run "
"`tailscale up` once before publishing"
) from exc
if result.returncode != 0:
details = "\n".join(
part.strip() for part in (result.stdout, result.stderr) if part.strip()
)
suffix = f"\n{details}" if details else ""
raise RuntimeError(f"command failed: {' '.join(command)}{suffix}")
return result
def _tailscale_dns_name(
tailscale_bin: str = "tailscale",
*,
runner: Callable[..., subprocess.CompletedProcess[str]] = subprocess.run,
) -> str:
result = _run_command([tailscale_bin, "status", "--json"], runner=runner)
try:
payload = json.loads(result.stdout)
dns_name = str(payload["Self"]["DNSName"]).strip().rstrip(".")
except (KeyError, TypeError, ValueError, json.JSONDecodeError) as exc:
raise RuntimeError(
"could not read this device's MagicDNS name from `tailscale status --json`; "
"confirm that Tailscale is connected and MagicDNS is enabled"
) from exc
if not dns_name or "." not in dns_name:
raise RuntimeError(f"Tailscale returned an invalid DNS name: {dns_name!r}")
return dns_name
def start_tailscale_funnel(
*,
local_port: int,
tailscale_bin: str = "tailscale",
runner: Callable[..., subprocess.CompletedProcess[str]] = subprocess.run,
) -> str:
"""Persistently publish a loopback HTTP service and return its HTTPS URL."""
if not 1 <= local_port <= 65535:
raise ValueError("local port must be between 1 and 65535")
dns_name = _tailscale_dns_name(tailscale_bin, runner=runner)
target = f"http://127.0.0.1:{local_port}"
_run_command(
[tailscale_bin, "funnel", "--bg", "--https=443", target],
runner=runner,
)
# A successful status call catches permission/configuration failures that
# some client versions report only after accepting the Funnel command.
_run_command(
[tailscale_bin, "funnel", "status", "--json"],
runner=runner,
)
return f"https://{dns_name}/forecast"
def publish_submission(
*,
model_id: str,
display_name: str,
code_url: str,
output_dir: Path,
local_port: int = 7860,
endpoint_url: str | None = None,
tailscale_bin: str = "tailscale",
wait_seconds: float = 600.0,
timeout: float = 90.0,
validator: Callable[..., dict[str, object]] = validate_endpoint,
funnel_starter: Callable[..., str] = start_tailscale_funnel,
) -> dict[str, object]:
"""Validate local/public routes and write a complete submission bundle."""
if not 1 <= local_port <= 65535:
raise ValueError("local port must be between 1 and 65535")
validate_submission_identity(
model_id=model_id,
display_name=display_name,
code_url=code_url,
)
if endpoint_url is not None:
# Fail fast on an invalid supplied route before making any requests.
build_metadata(
model_id=model_id,
display_name=display_name,
code_url=code_url,
endpoint_url=endpoint_url,
)
local_url = f"http://127.0.0.1:{local_port}/forecast"
local_receipt = validator(
endpoint_url=local_url,
model_id=model_id,
timeout=timeout,
require_https=False,
)
provider = "existing_https"
if endpoint_url is None:
endpoint_url = funnel_starter(
local_port=local_port,
tailscale_bin=tailscale_bin,
)
provider = "tailscale_funnel"
# The endpoint was unknown during the preflight when Funnel was selected,
# so validate the complete metadata before the public readiness wait.
metadata = build_metadata(
model_id=model_id,
display_name=display_name,
code_url=code_url,
endpoint_url=endpoint_url,
)
public_receipt = validator(
endpoint_url=endpoint_url,
model_id=model_id,
timeout=timeout,
wait_seconds=wait_seconds,
require_https=True,
)
receipt = {
"status": "ok",
"generated_at_utc": datetime.now(timezone.utc).isoformat(),
"provider": provider,
"endpoint_url": endpoint_url,
"local_validation": local_receipt,
"public_validation": public_receipt,
}
output_dir = output_dir.resolve()
output_dir.mkdir(parents=True, exist_ok=True)
metadata_path = output_dir / "community_model.yaml"
receipt_path = output_dir / "validator_receipt.json"
metadata_path.write_text(
yaml.safe_dump(metadata, sort_keys=False, allow_unicode=True),
encoding="utf-8",
)
receipt_path.write_text(
json.dumps(receipt, indent=2, ensure_ascii=False) + "\n",
encoding="utf-8",
)
return {
"status": "ok",
"endpoint_url": endpoint_url,
"metadata_path": str(metadata_path),
"receipt_path": str(receipt_path),
}
def parse_args() -> argparse.Namespace:
parser = argparse.ArgumentParser(description=__doc__)
subparsers = parser.add_subparsers(dest="command", required=True)
init_parser = subparsers.add_parser(
"init", help="copy the endpoint template into an editable service directory"
)
init_parser.add_argument("--output-dir", type=Path, default=Path("forecast-service"))
init_parser.add_argument("--force", action="store_true")
publish_parser = subparsers.add_parser(
"publish",
help="validate locally, publish through Funnel, and generate submission files",
)
publish_parser.add_argument("--model-id", required=True)
publish_parser.add_argument("--display-name", required=True)
publish_parser.add_argument("--code-url", required=True)
publish_parser.add_argument("--local-port", type=int, default=7860)
publish_parser.add_argument(
"--endpoint-url",
help="use an existing public HTTPS /forecast URL instead of Tailscale Funnel",
)
publish_parser.add_argument("--tailscale-bin", default="tailscale")
publish_parser.add_argument("--wait-seconds", type=float, default=600.0)
publish_parser.add_argument("--timeout", type=float, default=90.0)
publish_parser.add_argument(
"--output-dir", type=Path, default=Path("community-submission")
)
return parser.parse_args()
def main() -> int:
args = parse_args()
if args.command == "init":
copied = initialize_service(args.output_dir, force=args.force)
print(
json.dumps(
{
"status": "ok",
"service_dir": str(args.output_dir.resolve()),
"copied": [str(path) for path in copied],
"next": (
"Replace forecast_one, then run: docker compose -f "
f'"{args.output_dir.resolve() / "compose.yaml"}" '
"up -d --build"
),
},
indent=2,
)
)
return 0
result = publish_submission(
model_id=args.model_id,
display_name=args.display_name,
code_url=args.code_url,
output_dir=args.output_dir,
local_port=args.local_port,
endpoint_url=args.endpoint_url,
tailscale_bin=args.tailscale_bin,
wait_seconds=args.wait_seconds,
timeout=args.timeout,
)
print(json.dumps(result, indent=2))
return 0
if __name__ == "__main__":
try:
raise SystemExit(main())
except (RuntimeError, ValueError) as exc:
raise SystemExit(f"community model setup failed: {exc}") from exc