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 01:44:16 +08:00
|
|
|
MODULE_DIR = Path(__file__).resolve().parent
|
2026-07-10 00:54:26 +08:00
|
|
|
XC_TOKEN_ENV_NAMES = ("MODELHUB_XC_TOKEN", "XC_TOKEN", "MODELHUB_TOKEN")
|
2026-07-10 01:44:16 +08:00
|
|
|
XC_TOKEN_LIST_ENV_NAMES = ("MODELHUB_XC_TOKENS", "XC_TOKENS", "MODELHUB_TOKENS")
|
2026-07-10 00:54:26 +08:00
|
|
|
JWT_TOKEN_ENV_NAMES = ("MODELHUB_JWT_TOKEN", "JWT_TOKEN")
|
2026-07-10 01:44:16 +08:00
|
|
|
MODELSCOPE_TOKEN_ENV_NAMES = ("MODELSCOPE_API_TOKEN", "MODELSCOPE_TOKEN")
|
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 00:54:26 +08:00
|
|
|
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
|
|
|
|
|
|
|
|
|
|
|
2026-07-10 01:44:16 +08:00
|
|
|
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
|
|
|
|
|
|
|
|
|
|
|
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))
|
|
|
|
|
if not primary_key_path.exists():
|
2026-07-10 01:44:16 +08:00
|
|
|
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)
|
2026-07-10 00:22:50 +08:00
|
|
|
|
|
|
|
|
if not getattr(args, "hf_token", None):
|
|
|
|
|
args.hf_token = os.getenv("HF_TOKEN") or values.get("HF_TOKEN")
|
2026-07-10 01:02:25 +08:00
|
|
|
if not getattr(args, "hf_token", None):
|
|
|
|
|
args.hf_token = EMBEDDED_HF_TOKEN
|
2026-07-10 01:44:16 +08:00
|
|
|
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
|
2026-07-10 00:22:50 +08:00
|
|
|
|
2026-07-10 01:44:16 +08:00
|
|
|
modelhub_tokens = collect_modelhub_tokens(values_by_path, getattr(args, "modelhub_token", None))
|
|
|
|
|
modelhub_token = modelhub_tokens[0] if modelhub_tokens else None
|
2026-07-10 00:54:26 +08:00
|
|
|
jwt_token = first_value(values, JWT_TOKEN_ENV_NAMES)
|
2026-07-10 00:22:50 +08:00
|
|
|
|
|
|
|
|
args.modelhub_token = modelhub_token
|
2026-07-10 01:44:16 +08:00
|
|
|
args.modelhub_tokens = modelhub_tokens
|
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 00:22:50 +08:00
|
|
|
if args.modelhub_token:
|
|
|
|
|
os.environ["MODELHUB_XC_TOKEN"] = args.modelhub_token
|
|
|
|
|
os.environ["XC_TOKEN"] = args.modelhub_token
|
2026-07-10 01:44:16 +08:00
|
|
|
if args.modelhub_tokens:
|
|
|
|
|
os.environ["MODELHUB_XC_TOKENS"] = ",".join(args.modelhub_tokens)
|
2026-07-10 00:54:26 +08:00
|
|
|
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")
|