from __future__ import annotations import argparse import os from pathlib import Path from defaults import EMBEDDED_HF_TOKEN, EMBEDDED_MODELHUB_XC_TOKEN DEFAULT_KEY_PATH = Path("KEY.md") XC_TOKEN_ENV_NAMES = ("MODELHUB_XC_TOKEN", "XC_TOKEN", "MODELHUB_TOKEN") JWT_TOKEN_ENV_NAMES = ("MODELHUB_JWT_TOKEN", "JWT_TOKEN") def load_key_file(path: Path) -> dict[str, str]: if not path.exists(): return {} loaded: dict[str, str] = {} for raw_line in path.read_text(encoding="utf-8").splitlines(): line = raw_line.strip() if not line or line.startswith("#") or "=" not in line: continue key, value = line.split("=", 1) loaded[key.strip()] = value.strip() return loaded def first_value(values: dict[str, str], names: tuple[str, ...]) -> str | None: for name in names: value = os.getenv(name) or values.get(name) if value: return value.strip() return None def ensure_tokens(args: argparse.Namespace) -> None: primary_key_path = Path(getattr(args, "key_path", DEFAULT_KEY_PATH)) if not primary_key_path.exists(): parent_key = primary_key_path.parent.parent / primary_key_path.name if parent_key.exists(): primary_key_path = parent_key values = load_key_file(primary_key_path) if not getattr(args, "hf_token", None): args.hf_token = os.getenv("HF_TOKEN") or values.get("HF_TOKEN") if not getattr(args, "hf_token", None): args.hf_token = EMBEDDED_HF_TOKEN modelhub_token = getattr(args, "modelhub_token", None) if not modelhub_token: modelhub_token = first_value(values, XC_TOKEN_ENV_NAMES) if not modelhub_token: modelhub_token = EMBEDDED_MODELHUB_XC_TOKEN jwt_token = first_value(values, JWT_TOKEN_ENV_NAMES) args.modelhub_token = modelhub_token if args.hf_token: os.environ["HF_TOKEN"] = args.hf_token if args.modelhub_token: os.environ["MODELHUB_XC_TOKEN"] = args.modelhub_token os.environ["XC_TOKEN"] = args.modelhub_token if jwt_token: os.environ["MODELHUB_JWT_TOKEN"] = jwt_token if not args.modelhub_token and not jwt_token: raise ValueError("MODELHUB_XC_TOKEN/XC_TOKEN or MODELHUB_JWT_TOKEN/JWT_TOKEN is required")