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") DEFAULT_KEYS_PATH = Path("KEYS.md") MODULE_DIR = Path(__file__).resolve().parent MODELSCOPE_TOKEN_ENV_NAMES = ("MODELSCOPE_API_TOKEN", "MODELSCOPE_TOKEN") def _token_sort_key(key: str) -> tuple[int, str]: suffix = key.removeprefix("XC_TOKEN") if not suffix: return (0, key) if suffix.isdigit(): return (int(suffix), key) return (10_000, key) 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 load_key_files(*paths: Path) -> dict[str, str]: loaded: dict[str, str] = {} for path in paths: if path.exists(): loaded.update(load_key_file(path)) return loaded def load_modelhub_tokens(values: dict[str, str]) -> list[str]: tokens: list[tuple[str, str]] = [] seen: set[str] = set() for key, value in values.items(): if key == "MODELHUB_XC_TOKEN" or key.startswith("XC_TOKEN"): token = value.strip() if token and token not in seen: seen.add(token) tokens.append((key, token)) tokens.sort(key=lambda item: _token_sort_key(item[0])) return [token for _, token in tokens] def _split_token_list(value: str | None) -> list[str]: if not value: return [] tokens: list[str] = [] for raw in value.replace(",", "\n").replace(";", "\n").splitlines(): token = raw.strip() if token: tokens.append(token) return tokens def _add_token(tokens: list[str], token: str | None) -> None: for item in _split_token_list(token): if item and item not in tokens: tokens.append(item) def ensure_tokens(args: argparse.Namespace) -> None: primary_key_path = Path(getattr(args, "key_path", DEFAULT_KEY_PATH)) supplemental_key_path = primary_key_path.with_name(DEFAULT_KEYS_PATH.name) if not primary_key_path.exists(): for candidate in ( MODULE_DIR / primary_key_path.name, primary_key_path.parent.parent / primary_key_path.name, MODULE_DIR.parent / primary_key_path.name, ): if candidate.exists(): primary_key_path = candidate supplemental_key_path = candidate.with_name(DEFAULT_KEYS_PATH.name) break values = load_key_files(primary_key_path, supplemental_key_path) if not getattr(args, "hf_token", None): args.hf_token = values.get("HF_TOKEN") or EMBEDDED_HF_TOKEN if not getattr(args, "modelscope_token", None): args.modelscope_token = ( os.getenv("MODELSCOPE_API_TOKEN") or os.getenv("MODELSCOPE_TOKEN") or values.get("MODELSCOPE_API_TOKEN") or values.get("MODELSCOPE_TOKEN") or EMBEDDED_MODELSCOPE_TOKEN ) tokens = load_modelhub_tokens(values) env_token_items: list[tuple[str, str]] = [] for key, value in os.environ.items(): if key == "MODELHUB_XC_TOKEN" or key.startswith("XC_TOKEN"): token = value.strip() if token: env_token_items.append((key, token)) env_token_items.sort(key=lambda item: _token_sort_key(item[0])) env_tokens = [token for _, token in env_token_items] modelhub_token = getattr(args, "modelhub_token", None) if modelhub_token and modelhub_token not in env_tokens: env_tokens.insert(0, modelhub_token) if not env_tokens: env_tokens = tokens else: for token in tokens: if token not in env_tokens: env_tokens.append(token) for env_name in ("MODELHUB_XC_TOKENS", "XC_TOKENS", "MODELHUB_TOKENS"): for token in _split_token_list(os.getenv(env_name)): _add_token(env_tokens, token) for token in EMBEDDED_MODELHUB_XC_TOKENS: _add_token(env_tokens, token) args.modelhub_tokens = env_tokens args.modelhub_token = env_tokens[0] if env_tokens else None 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 for index, token in enumerate(args.modelhub_tokens, start=1): env_name = "XC_TOKEN" if index == 1 else f"XC_TOKEN{index}" os.environ[env_name] = token if args.modelhub_token: os.environ["MODELHUB_XC_TOKEN"] = args.modelhub_token os.environ["MODELHUB_XC_TOKENS"] = ",".join(args.modelhub_tokens) if not args.modelhub_token: raise ValueError("XC_TOKEN is required, either via environment or KEY.md")