Files
xc_validation_strategy_vllm…/main.py
2026-07-25 15:30:17 +08:00

249 lines
8.3 KiB
Python

"""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()