"""Stop zhoukaile's waiting ModelHub XC validation tasks. The service performs the stop requests once at startup and then stays alive so the strategy platform can probe it through ``/health`` and inspect ``/status``. """ import json import os import signal import threading import time from datetime import datetime, timezone from http.server import BaseHTTPRequestHandler, ThreadingHTTPServer from typing import Any import requests BASE_URL = os.environ.get("BASE_URL", "https://modelhub.org.cn").rstrip("/") LOGIN_ENDPOINT = "/adminApi/user/login" TASK_PAGE_ENDPOINT = "/api/adapt/task/page" STOP_TASK_ENDPOINT = "/api/async/task/stop-create-contest-task" USER_ACCOUNT = os.environ.get("USER_ACCOUNT", "zhoukaile") USER_PASSWORD = os.environ.get("USER_PASSWORD", "") XC_TOKEN = os.environ.get("XC_TOKEN", "") STRATEGY_ID = os.environ.get("STRATEGY_ID", "") HTTP_HOST = "0.0.0.0" HTTP_PORT = int(os.environ.get("PORT", "8080")) BATCH_SIZE = int(os.environ.get("BATCH_SIZE", "50")) MAX_RETRIES = int(os.environ.get("MAX_RETRIES", "3")) REQUEST_TIMEOUT = int(os.environ.get("REQUEST_TIMEOUT", "30")) PAGE_SIZE = 100 STOPPABLE_STATUS = "waiting" _shutdown = threading.Event() _state: dict[str, Any] = { "strategy_id": STRATEGY_ID, "account": USER_ACCOUNT, "phase": "starting", # starting | stopping | done | partial_failure | error "total": 0, "stopped": 0, "failed": 0, "failed_task_ids": [], "started_at": None, "finished_at": None, "error": None, } def _now() -> str: return datetime.now(timezone.utc).isoformat() def _task_ids_from_environment() -> list[int] | None: """Return an explicit TASK_IDS override, if one was supplied.""" raw_task_ids = os.environ.get("TASK_IDS", "").strip() if not raw_task_ids: return None task_ids: list[int] = [] for value in raw_task_ids.split(","): value = value.strip() if not value: continue try: task_ids.append(int(value)) except ValueError as exc: raise ValueError(f"TASK_IDS contains a non-numeric task ID: {value!r}") from exc return list(dict.fromkeys(task_ids)) def _login() -> str: """Authenticate as the task owner and return a short-lived bearer token.""" if not USER_PASSWORD: raise RuntimeError("USER_PASSWORD is required to query zhoukaile's current task IDs") response = requests.post( f"{BASE_URL}{LOGIN_ENDPOINT}", headers={"Content-Type": "application/json"}, json={"userAccount": USER_ACCOUNT, "userPassword": USER_PASSWORD}, timeout=REQUEST_TIMEOUT, ) response.raise_for_status() result = response.json() token = result.get("data", {}).get("token") if result.get("code") != 0 or not isinstance(token, str) or not token: raise RuntimeError(f"Task query login failed: {result.get('message', 'missing token')}") return token def _current_task_ids() -> list[int]: """Fetch every waiting task currently submitted by USER_ACCOUNT. This intentionally queries the user-facing task page at run time. A static list becomes obsolete as soon as zhoukaile submits more validation tasks. """ bearer_token = _login() headers = {"Authorization": f"Bearer {bearer_token}"} task_ids: list[int] = [] page = 1 while True: response = requests.get( f"{BASE_URL}{TASK_PAGE_ENDPOINT}", headers=headers, params={"current": page, "pageSize": PAGE_SIZE}, timeout=REQUEST_TIMEOUT, ) response.raise_for_status() result = response.json() if result.get("code") != 0: raise RuntimeError(f"Task query failed: {result.get('message', 'unknown error')}") data = result.get("data") or {} records = data.get("records") or [] for record in records: if str(record.get("status", "")).lower() != STOPPABLE_STATUS: continue try: task_ids.append(int(record["taskId"])) except (KeyError, TypeError, ValueError): print(f"[query] skipping record without a numeric taskId: {record}", flush=True) pages = int(data.get("pages") or 0) if page >= pages: break page += 1 return list(dict.fromkeys(task_ids)) class Handler(BaseHTTPRequestHandler): def do_GET(self) -> None: # noqa: N802 - BaseHTTPRequestHandler API if self.path == "/health": self._json({"status": "ok"}) elif self.path == "/status": self._json(_state) else: self._json({"error": "not found"}, 404) def _json(self, body: dict[str, Any], status: int = 200) -> None: payload = json.dumps(body, ensure_ascii=False).encode("utf-8") self.send_response(status) self.send_header("Content-Type", "application/json; charset=utf-8") self.send_header("Content-Length", str(len(payload))) self.end_headers() self.wfile.write(payload) def log_message(self, fmt: str, *args: object) -> None: print(f"[http] {self.address_string()} {fmt % args}", flush=True) def _run_http() -> None: server = ThreadingHTTPServer((HTTP_HOST, HTTP_PORT), Handler) server.timeout = 1 print(f"[http] listening on {HTTP_HOST}:{HTTP_PORT}", flush=True) while not _shutdown.is_set(): server.handle_request() server.server_close() def _stop_batch(task_ids: list[int]) -> bool: headers = {"Content-Type": "application/json", "xc-Token": XC_TOKEN} url = f"{BASE_URL}{STOP_TASK_ENDPOINT}" for attempt in range(1, MAX_RETRIES + 1): try: response = requests.put( url, headers=headers, json={"taskIds": task_ids}, timeout=REQUEST_TIMEOUT, ) try: result = response.json() except ValueError: result = {"message": response.text[:500]} if response.ok and result.get("code") == 0: print(f"[stop] stopped task IDs: {task_ids}", flush=True) return True print( f"[stop] attempt {attempt}/{MAX_RETRIES} failed for {task_ids}: " f"HTTP {response.status_code}, {result}", flush=True, ) except requests.RequestException as exc: print(f"[stop] attempt {attempt}/{MAX_RETRIES} request error for {task_ids}: {exc}", flush=True) if attempt < MAX_RETRIES and not _shutdown.wait(attempt): continue if _shutdown.is_set(): break return False def _run_worker() -> None: _state["started_at"] = _now() _state["phase"] = "stopping" try: if not XC_TOKEN: raise RuntimeError("XC_TOKEN is required and must be configured in the strategy environment") task_ids = _task_ids_from_environment() or _current_task_ids() if not task_ids: print("[query] no active validation tasks found", flush=True) _state["phase"] = "done" return _state["total"] = len(task_ids) for start in range(0, len(task_ids), BATCH_SIZE): if _shutdown.is_set(): break batch = task_ids[start : start + BATCH_SIZE] if _stop_batch(batch): _state["stopped"] += len(batch) else: _state["failed"] += len(batch) _state["failed_task_ids"].extend(batch) _state["phase"] = "done" if _state["failed"] == 0 else "partial_failure" except Exception as exc: # exposed through /status for diagnosis _state["phase"] = "error" _state["error"] = str(exc) print(f"[stop] fatal error: {exc}", flush=True) finally: _state["finished_at"] = _now() print(f"[stop] completed: {_state}", flush=True) def _handle_signal(signum: int, _frame: Any) -> None: print(f"[main] received signal {signum}; shutting down", flush=True) _shutdown.set() def main() -> None: signal.signal(signal.SIGTERM, _handle_signal) signal.signal(signal.SIGINT, _handle_signal) http_thread = threading.Thread(target=_run_http, daemon=False) http_thread.start() threading.Thread(target=_run_worker, daemon=True).start() _shutdown.wait() http_thread.join(timeout=5) if __name__ == "__main__": main()