2026-07-10 00:22:50 +08:00
|
|
|
from __future__ import annotations
|
|
|
|
|
|
|
|
|
|
import argparse
|
|
|
|
|
import os
|
|
|
|
|
from pathlib import Path
|
|
|
|
|
|
2026-07-10 01:44:16 +08:00
|
|
|
from defaults import EMBEDDED_HF_TOKEN, EMBEDDED_MODELHUB_XC_TOKENS, EMBEDDED_MODELSCOPE_TOKEN
|
2026-07-10 00:54:26 +08:00
|
|
|
|
2026-07-10 00:22:50 +08:00
|
|
|
|
|
|
|
|
DEFAULT_KEY_PATH = Path("KEY.md")
|
2026-07-10 02:02:08 +08:00
|
|
|
DEFAULT_KEYS_PATH = Path("KEYS.md")
|
2026-07-10 01:44:16 +08:00
|
|
|
MODULE_DIR = Path(__file__).resolve().parent
|
|
|
|
|
MODELSCOPE_TOKEN_ENV_NAMES = ("MODELSCOPE_API_TOKEN", "MODELSCOPE_TOKEN")
|
2026-07-10 00:22:50 +08:00
|
|
|
|
|
|
|
|
|
2026-07-10 02:02:08 +08:00
|
|
|
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)
|
|
|
|
|
|
|
|
|
|
|
2026-07-10 00:22:50 +08:00
|
|
|
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
|
|
|
|
|
|
|
|
|
|
|
2026-07-10 02:02:08 +08:00
|
|
|
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]
|
2026-07-10 00:54:26 +08:00
|
|
|
|
|
|
|
|
|
2026-07-10 02:02:08 +08:00
|
|
|
def _split_token_list(value: str | None) -> list[str]:
|
2026-07-10 01:44:16 +08:00
|
|
|
if not value:
|
|
|
|
|
return []
|
|
|
|
|
tokens: list[str] = []
|
2026-07-10 02:02:08 +08:00
|
|
|
for raw in value.replace(",", "\n").replace(";", "\n").splitlines():
|
|
|
|
|
token = raw.strip()
|
2026-07-10 01:44:16 +08:00
|
|
|
if token:
|
|
|
|
|
tokens.append(token)
|
|
|
|
|
return tokens
|
|
|
|
|
|
|
|
|
|
|
2026-07-10 02:02:08 +08:00
|
|
|
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)
|
2026-07-10 01:44:16 +08:00
|
|
|
|
|
|
|
|
|
2026-07-10 00:22:50 +08:00
|
|
|
def ensure_tokens(args: argparse.Namespace) -> None:
|
|
|
|
|
primary_key_path = Path(getattr(args, "key_path", DEFAULT_KEY_PATH))
|
2026-07-10 02:02:08 +08:00
|
|
|
supplemental_key_path = primary_key_path.with_name(DEFAULT_KEYS_PATH.name)
|
2026-07-10 00:22:50 +08:00
|
|
|
if not primary_key_path.exists():
|
2026-07-10 02:02:08 +08:00
|
|
|
for candidate in (
|
2026-07-10 01:44:16 +08:00
|
|
|
MODULE_DIR / primary_key_path.name,
|
|
|
|
|
primary_key_path.parent.parent / primary_key_path.name,
|
|
|
|
|
MODULE_DIR.parent / primary_key_path.name,
|
2026-07-10 02:02:08 +08:00
|
|
|
):
|
|
|
|
|
if candidate.exists():
|
|
|
|
|
primary_key_path = candidate
|
|
|
|
|
supplemental_key_path = candidate.with_name(DEFAULT_KEYS_PATH.name)
|
2026-07-10 01:44:16 +08:00
|
|
|
break
|
2026-07-10 02:02:08 +08:00
|
|
|
values = load_key_files(primary_key_path, supplemental_key_path)
|
2026-07-10 00:22:50 +08:00
|
|
|
|
|
|
|
|
if not getattr(args, "hf_token", None):
|
2026-07-10 02:02:08 +08:00
|
|
|
args.hf_token = values.get("HF_TOKEN") or EMBEDDED_HF_TOKEN
|
2026-07-10 01:44:16 +08:00
|
|
|
if not getattr(args, "modelscope_token", None):
|
2026-07-10 02:02:08 +08:00
|
|
|
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)
|
2026-07-10 00:22:50 +08:00
|
|
|
|
2026-07-10 02:02:08 +08:00
|
|
|
args.modelhub_tokens = env_tokens
|
|
|
|
|
args.modelhub_token = env_tokens[0] if env_tokens else None
|
2026-07-10 00:22:50 +08:00
|
|
|
|
|
|
|
|
if args.hf_token:
|
|
|
|
|
os.environ["HF_TOKEN"] = args.hf_token
|
2026-07-10 01:44:16 +08:00
|
|
|
if args.modelscope_token:
|
|
|
|
|
os.environ["MODELSCOPE_API_TOKEN"] = args.modelscope_token
|
|
|
|
|
os.environ["MODELSCOPE_TOKEN"] = args.modelscope_token
|
2026-07-10 02:02:08 +08:00
|
|
|
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
|
2026-07-10 00:22:50 +08:00
|
|
|
if args.modelhub_token:
|
|
|
|
|
os.environ["MODELHUB_XC_TOKEN"] = args.modelhub_token
|
2026-07-10 01:44:16 +08:00
|
|
|
os.environ["MODELHUB_XC_TOKENS"] = ",".join(args.modelhub_tokens)
|
2026-07-10 02:02:08 +08:00
|
|
|
if not args.modelhub_token:
|
|
|
|
|
raise ValueError("XC_TOKEN is required, either via environment or KEY.md")
|