from __future__ import annotations import argparse import os from pathlib import Path from defaults import EMBEDDED_HF_TOKEN, EMBEDDED_MODELHUB_XC_TOKENS, EMBEDDED_MODELSCOPE_TOKEN DEFAULT_KEY_PATH = Path("KEY.md") MODULE_DIR = Path(__file__).resolve().parent XC_TOKEN_ENV_NAMES = ("MODELHUB_XC_TOKEN", "XC_TOKEN", "MODELHUB_TOKEN") XC_TOKEN_LIST_ENV_NAMES = ("MODELHUB_XC_TOKENS", "XC_TOKENS", "MODELHUB_TOKENS") JWT_TOKEN_ENV_NAMES = ("MODELHUB_JWT_TOKEN", "JWT_TOKEN") MODELSCOPE_TOKEN_ENV_NAMES = ("MODELSCOPE_API_TOKEN", "MODELSCOPE_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 split_token_list(value: str | None) -> list[str]: if not value: return [] tokens: list[str] = [] for normalized in value.replace(",", "\n").replace(";", "\n").splitlines(): token = normalized.strip() if token: tokens.append(token) return tokens def discover_key_paths(primary_path: Path) -> list[Path]: candidates: list[Path] = [] search_dirs = [primary_path.parent, primary_path.parent.parent] for path in [primary_path, primary_path.parent / "KEYS.md"]: if path not in candidates: candidates.append(path) for directory in search_dirs: if not directory.exists(): continue for path in sorted(directory.glob("KEYS*")): if path.is_file() and path not in candidates: candidates.append(path) return candidates def collect_modelhub_tokens(values_by_path: list[dict[str, str]], explicit_token: str | None) -> list[str]: tokens: list[str] = [] def add(value: str | None) -> None: for token in split_token_list(value): if token and token not in tokens: tokens.append(token) add(explicit_token) for env_name in XC_TOKEN_LIST_ENV_NAMES: add(os.getenv(env_name)) for env_name in XC_TOKEN_ENV_NAMES: add(os.getenv(env_name)) for values in values_by_path: for key, value in values.items(): normalized = key.strip().upper() if normalized in XC_TOKEN_ENV_NAMES or normalized.startswith("XC_TOKEN") or normalized.startswith("MODELHUB_XC_TOKEN"): add(value) for token in EMBEDDED_MODELHUB_XC_TOKENS: add(token) return tokens def ensure_tokens(args: argparse.Namespace) -> None: primary_key_path = Path(getattr(args, "key_path", DEFAULT_KEY_PATH)) if not primary_key_path.exists(): fallback_paths = [ MODULE_DIR / primary_key_path.name, primary_key_path.parent.parent / primary_key_path.name, MODULE_DIR.parent / primary_key_path.name, ] for fallback_path in fallback_paths: if fallback_path.exists(): primary_key_path = fallback_path break values_by_path = [load_key_file(path) for path in discover_key_paths(primary_key_path) if path.exists()] values: dict[str, str] = {} for loaded in values_by_path: values.update(loaded) 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 if not getattr(args, "modelscope_token", None): args.modelscope_token = first_value(values, MODELSCOPE_TOKEN_ENV_NAMES) if not getattr(args, "modelscope_token", None): args.modelscope_token = EMBEDDED_MODELSCOPE_TOKEN modelhub_tokens = collect_modelhub_tokens(values_by_path, getattr(args, "modelhub_token", None)) modelhub_token = modelhub_tokens[0] if modelhub_tokens else None jwt_token = first_value(values, JWT_TOKEN_ENV_NAMES) args.modelhub_token = modelhub_token args.modelhub_tokens = modelhub_tokens if args.hf_token: os.environ["HF_TOKEN"] = args.hf_token if args.modelscope_token: os.environ["MODELSCOPE_API_TOKEN"] = args.modelscope_token os.environ["MODELSCOPE_TOKEN"] = args.modelscope_token if args.modelhub_token: os.environ["MODELHUB_XC_TOKEN"] = args.modelhub_token os.environ["XC_TOKEN"] = args.modelhub_token if args.modelhub_tokens: os.environ["MODELHUB_XC_TOKENS"] = ",".join(args.modelhub_tokens) 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")