Files
TradeBot/crypto_spot_bot/training_coordination.py
T

634 lines
25 KiB
Python

from __future__ import annotations
import base64
import hashlib
import hmac
import json
import os
import re
import secrets
import shutil
import uuid
from datetime import UTC
from datetime import datetime
from datetime import timedelta
from pathlib import Path
from threading import Lock
from typing import Any
ALLOWED_TRAINING_ARTIFACTS = {
"lstm_forecaster.json",
"torch_retrain_guard.json",
"torch_threshold_calibration.json",
}
RUNNING_LEASE_TIMEOUT = timedelta(minutes=10)
ONLINE_WINDOW = timedelta(minutes=3)
MAX_JOB_ATTEMPTS = 3
MAX_ARTIFACT_CHUNK_BYTES = 1024 * 1024
# Keep uploads bounded while leaving room for explicitly requested per-symbol bundles.
MAX_ARTIFACT_BYTES = 256 * 1024 * 1024
MAX_ARTIFACT_CHUNKS = 1024
REQUIRED_MODEL_BUNDLE = set(ALLOWED_TRAINING_ARTIFACTS)
class TrainingCoordinator:
def __init__(self, runtime_dir: Path) -> None:
self.runtime_dir = runtime_dir
self.state_path = runtime_dir / "training_coordination.json"
self.upload_root = runtime_dir / ".training_uploads"
self._lock = Lock()
def status(self) -> dict[str, Any]:
with self._lock:
state = self._load_state()
self._expire_stale_jobs(state)
self._save_state(state)
return self._public_status(state)
def request_retrain(self, payload: dict[str, Any] | None = None) -> dict[str, Any]:
payload = payload or {}
with self._lock:
state = self._load_state()
self._expire_stale_jobs(state)
existing = self._active_job(state)
if existing is not None:
self._save_state(state)
return {
"queued": False,
"reason": "active_job_exists",
"job": self._public_job(existing),
"status": self._public_status(state),
}
now = _now()
job = {
"id": str(uuid.uuid4()),
"status": "pending",
"requested_at": now,
"requested_by": str(payload.get("source") or "api"),
"parameters": _safe_parameters(payload.get("parameters")),
"message": "",
"artifacts": [],
"attempts": 0,
}
state.setdefault("jobs", []).append(job)
self._trim_jobs(state)
self._save_state(state)
return {
"queued": True,
"job": self._public_job(job),
"status": self._public_status(state),
}
def heartbeat(self, payload: dict[str, Any] | None = None) -> dict[str, Any]:
payload = payload or {}
with self._lock:
state = self._load_state()
worker = self._worker_from_payload(payload)
state["worker"] = worker
self._save_state(state)
return {"ok": True, "worker": worker, "status": self._public_status(state)}
def claim(self, payload: dict[str, Any] | None = None) -> dict[str, Any]:
payload = payload or {}
with self._lock:
state = self._load_state()
self._expire_stale_jobs(state)
worker = self._worker_from_payload(payload)
state["worker"] = worker
job = self._oldest_pending_job(state)
if job is None:
self._save_state(state)
return {"claimed": False, "job": None, "status": self._public_status(state)}
now = _now()
lease_token = secrets.token_urlsafe(32)
job["status"] = "running"
job["claimed_at"] = now
job["updated_at"] = now
job["claimed_by"] = worker["id"]
job["worker"] = worker
job["lease_token"] = lease_token
job["attempts"] = int(job.get("attempts", 0)) + 1
self._save_state(state)
return {
"claimed": True,
"job": self._public_job(job),
"lease_token": lease_token,
"status": self._public_status(state),
}
def save_artifact_chunk(self, job_id: str, payload: dict[str, Any]) -> dict[str, Any]:
job_id = _valid_job_id(job_id)
name = Path(str(payload.get("name") or "")).name
if name not in ALLOWED_TRAINING_ARTIFACTS:
raise ValueError(f"artifact is not allowed: {name}")
index = int(payload.get("index", -1))
total = int(payload.get("total", 0))
sha256 = str(payload.get("sha256") or "").strip().lower()
if index < 0 or total <= 0 or index >= total:
raise ValueError("invalid artifact chunk index")
if total > MAX_ARTIFACT_CHUNKS:
raise ValueError("artifact has too many chunks")
if not re.fullmatch(r"[0-9a-f]{64}", sha256):
raise ValueError("artifact sha256 is invalid")
try:
chunk = base64.b64decode(str(payload.get("data_base64") or ""), validate=True)
except (ValueError, TypeError) as exc:
raise ValueError("invalid artifact chunk payload") from exc
if not chunk or len(chunk) > MAX_ARTIFACT_CHUNK_BYTES:
raise ValueError("artifact chunk size is invalid")
with self._lock:
state = self._load_state()
job = self._job_by_id(state, job_id)
if job is None:
raise ValueError(f"training job not found: {job_id}")
if job.get("status") != "running" or not job.get("claimed_by"):
raise ValueError("training job is not claimed and running")
self._require_lease(job, payload)
uploads = job.setdefault("uploads", {})
upload = uploads.setdefault(name, {"sha256": sha256, "total": total})
if upload.get("sha256") != sha256 or int(upload.get("total", 0)) != total:
raise ValueError("artifact upload metadata changed during upload")
chunk_dir = self.upload_root / job_id / "chunks" / name
chunk_dir.mkdir(parents=True, exist_ok=True)
(chunk_dir / f"{index:06d}.part").write_bytes(chunk)
received = sum(1 for part in range(total) if (chunk_dir / f"{part:06d}.part").is_file())
if received < total:
upload["received"] = received
job["updated_at"] = _now()
self._save_state(state)
return {"complete": False, "received": received, "total": total}
ready_dir = self.upload_root / job_id / "ready"
ready_dir.mkdir(parents=True, exist_ok=True)
target_tmp = ready_dir / f".{name}.tmp"
digest = hashlib.sha256()
size = 0
with target_tmp.open("wb") as output:
for part in range(total):
data = (chunk_dir / f"{part:06d}.part").read_bytes()
size += len(data)
if size > MAX_ARTIFACT_BYTES:
target_tmp.unlink(missing_ok=True)
raise ValueError("artifact exceeds maximum size")
digest.update(data)
output.write(data)
if digest.hexdigest().lower() != sha256:
target_tmp.unlink(missing_ok=True)
raise ValueError("artifact sha256 mismatch")
target = ready_dir / name
os.replace(target_tmp, target)
_remove_tree(chunk_dir)
artifacts = job.setdefault("artifacts", [])
artifacts = [item for item in artifacts if item.get("name") != name]
artifacts.append(
{"name": name, "sha256": sha256, "size": size, "staged_at": _now()}
)
job["artifacts"] = artifacts
job["updated_at"] = _now()
upload["received"] = total
upload["complete"] = True
self._save_state(state)
return {"complete": True, "staged": True, "name": name, "sha256": sha256}
def progress(self, job_id: str, payload: dict[str, Any] | None = None) -> dict[str, Any]:
payload = payload or {}
job_id = _valid_job_id(job_id)
with self._lock:
state = self._load_state()
job = self._job_by_id(state, job_id)
if job is None:
raise ValueError(f"training job not found: {job_id}")
if job.get("status") != "running" or not job.get("claimed_by"):
raise ValueError("training job is not claimed and running")
self._require_lease(job, payload)
if isinstance(payload.get("worker"), dict):
state["worker"] = self._worker_from_payload(payload["worker"])
job["status"] = "running"
job["phase"] = str(payload.get("phase") or job.get("phase") or "running")[:80]
job["message"] = str(payload.get("message") or job.get("message") or "")[:2000]
job["progress_percent"] = _coerce_percent(payload.get("progress_percent"), job.get("progress_percent", 0))
job["updated_at"] = _now()
if isinstance(payload.get("details"), dict):
job["details"] = payload["details"]
self._save_state(state)
return {
"ok": True,
"job": self._public_job(job),
"status": self._public_status(state),
}
def complete(self, job_id: str, payload: dict[str, Any] | None = None) -> dict[str, Any]:
payload = payload or {}
job_id = _valid_job_id(job_id)
with self._lock:
state = self._load_state()
job = self._job_by_id(state, job_id)
if job is None:
raise ValueError(f"training job not found: {job_id}")
if job.get("status") != "running" or not job.get("claimed_by"):
raise ValueError("training job is not claimed and running")
self._require_lease(job, payload)
success = bool(payload.get("success", payload.get("status") == "completed"))
if success and job.get("artifacts"):
promoted = self._validate_and_promote(job_id, job)
job["promoted_artifacts"] = promoted
job["status"] = "completed" if success else "failed"
job["phase"] = "completed" if success else "failed"
job["progress_percent"] = 100 if success else _coerce_percent(payload.get("progress_percent"), job.get("progress_percent", 0))
job["completed_at"] = _now()
job["message"] = str(payload.get("message") or "")
if isinstance(payload.get("summary"), dict):
job["summary"] = payload["summary"]
if isinstance(payload["summary"].get("accepted"), bool):
job["model_decision"] = (
"accepted" if payload["summary"]["accepted"] else "rejected"
)
job.pop("lease_token", None)
self._save_state(state)
return {
"ok": True,
"job": self._public_job(job),
"status": self._public_status(state),
}
def _validate_and_promote(self, job_id: str, job: dict[str, Any]) -> list[dict[str, Any]]:
ready_dir = self.upload_root / job_id / "ready"
staged = {path.name for path in ready_dir.iterdir() if path.is_file()} if ready_dir.is_dir() else set()
missing = REQUIRED_MODEL_BUNDLE - staged
if missing:
raise ValueError("training bundle is incomplete: " + ", ".join(sorted(missing)))
model = _read_json(ready_dir / "lstm_forecaster.json")
guard = _read_json(ready_dir / "torch_retrain_guard.json")
calibration = _read_json(ready_dir / "torch_threshold_calibration.json")
if model.get("type") != "pytorch_recurrent_forecaster":
raise ValueError("candidate model type is invalid")
symbols = model.get("symbols")
if not isinstance(symbols, dict) or not symbols:
raise ValueError("candidate model has no symbol models")
_validate_symbol_models(symbols)
model_sha256 = hashlib.sha256((ready_dir / "lstm_forecaster.json").read_bytes()).hexdigest()
if calibration.get("artifact_sha256") != model_sha256:
raise ValueError("candidate calibration is not bound to the uploaded model")
if not bool(guard.get("accepted")):
raise ValueError("candidate retrain guard did not accept the model")
if guard.get("candidate_artifact_sha256") != model_sha256:
raise ValueError("candidate guard is not bound to the uploaded model")
validation = calibration.get("validation")
if not isinstance(validation, dict) or not _validation_passed(validation):
raise ValueError("candidate quality gate did not pass")
if validation.get("protocol") != "untouched_model_holdout_with_threshold_walk_forward":
raise ValueError("candidate validation protocol is not an untouched holdout")
self.runtime_dir.mkdir(parents=True, exist_ok=True)
backup_dir = self.runtime_dir / ".model_backups" / f"{_compact_now()}-{job_id}"
backup_dir.mkdir(parents=True, exist_ok=True)
for name in sorted(REQUIRED_MODEL_BUNDLE):
current = self.runtime_dir / name
if current.is_file():
shutil.copy2(current, backup_dir / name)
promoted: list[dict[str, Any]] = []
artifact_rows = {
str(item.get("name")): item
for item in job.get("artifacts", [])
if isinstance(item, dict)
}
for name in sorted(REQUIRED_MODEL_BUNDLE):
staged_path = ready_dir / name
target_tmp = self.runtime_dir / f".{name}.{job_id}.promote"
shutil.copy2(staged_path, target_tmp)
os.replace(target_tmp, self.runtime_dir / name)
row = artifact_rows.get(name, {})
promoted.append(
{
"name": name,
"sha256": row.get("sha256", ""),
"promoted_at": _now(),
}
)
_remove_tree(self.upload_root / job_id)
return promoted
def _load_state(self) -> dict[str, Any]:
try:
data = json.loads(self.state_path.read_text(encoding="utf-8"))
except (OSError, json.JSONDecodeError):
data = {}
if not isinstance(data, dict):
data = {}
data.setdefault("jobs", [])
return data
def _save_state(self, state: dict[str, Any]) -> None:
self.runtime_dir.mkdir(parents=True, exist_ok=True)
tmp = self.state_path.with_suffix(".tmp")
tmp.write_text(json.dumps(state, ensure_ascii=False, indent=2) + "\n", encoding="utf-8")
os.replace(tmp, self.state_path)
def _worker_from_payload(self, payload: dict[str, Any]) -> dict[str, Any]:
worker_id = str(payload.get("worker_id") or payload.get("id") or "windows-training-host").strip()
return {
"id": worker_id,
"name": str(payload.get("name") or worker_id).strip(),
"path": str(payload.get("path") or "").strip(),
"version": str(payload.get("version") or "1"),
"last_seen_at": _now(),
}
def _public_status(self, state: dict[str, Any]) -> dict[str, Any]:
worker = state.get("worker") if isinstance(state.get("worker"), dict) else {}
last_seen = _parse_time(str(worker.get("last_seen_at") or ""))
active = self._active_job(state)
latest = _latest_job(state)
recently_seen = bool(last_seen and datetime.now(UTC) - last_seen <= ONLINE_WINDOW)
agent_busy = bool(
active
and active.get("status") == "running"
and worker
and active.get("claimed_by") == worker.get("id")
)
return {
"available": True,
"agent_online": recently_seen or agent_busy,
"agent_recently_seen": recently_seen,
"agent_busy": agent_busy,
"worker": worker,
"active_job": self._public_job(active),
"latest_job": self._public_job(latest),
"pending_jobs": sum(1 for job in state.get("jobs", []) if job.get("status") == "pending"),
}
@staticmethod
def _public_job(job: dict[str, Any] | None) -> dict[str, Any] | None:
if job is None:
return None
public = dict(job)
public.pop("lease_token", None)
public.pop("uploads", None)
return public
@staticmethod
def _require_lease(job: dict[str, Any], payload: dict[str, Any]) -> None:
expected = str(job.get("lease_token") or "")
supplied = str(payload.get("lease_token") or "")
if not expected or not supplied or not hmac.compare_digest(expected, supplied):
raise ValueError("training job lease is invalid or expired")
def _active_job(self, state: dict[str, Any]) -> dict[str, Any] | None:
for job in reversed(state.get("jobs", [])):
if job.get("status") in {"pending", "running"}:
return job
return None
def _oldest_pending_job(self, state: dict[str, Any]) -> dict[str, Any] | None:
for job in state.get("jobs", []):
if job.get("status") == "pending":
return job
return None
def _job_by_id(self, state: dict[str, Any], job_id: str) -> dict[str, Any] | None:
for job in state.get("jobs", []):
if job.get("id") == job_id:
return job
return None
def _expire_stale_jobs(self, state: dict[str, Any]) -> None:
now = datetime.now(UTC)
for job in state.get("jobs", []):
if job.get("status") != "running":
continue
lease_updated_at = _parse_time(
str(job.get("updated_at") or job.get("claimed_at") or "")
)
if not lease_updated_at or now - lease_updated_at <= RUNNING_LEASE_TIMEOUT:
continue
job_id = str(job.get("id") or "")
if job_id:
_remove_tree(self.upload_root / job_id)
job.pop("lease_token", None)
job.pop("uploads", None)
attempts = int(job.get("attempts", 0))
if attempts < MAX_JOB_ATTEMPTS:
job["status"] = "pending"
job["phase"] = "queued"
job["progress_percent"] = 0
job["message"] = "training worker lease expired; queued for retry"
job["retry_queued_at"] = _now()
for key in ("claimed_at", "claimed_by", "worker", "updated_at"):
job.pop(key, None)
else:
job["status"] = "failed"
job["phase"] = "failed"
job["completed_at"] = _now()
job["message"] = "training worker lease expired after maximum retries"
def _trim_jobs(self, state: dict[str, Any]) -> None:
jobs = state.get("jobs", [])
if isinstance(jobs, list) and len(jobs) > 30:
state["jobs"] = jobs[-30:]
def _safe_parameters(value: Any) -> dict[str, Any]:
if not isinstance(value, dict):
return {}
allowed = {
"symbols",
"limit",
"lookbacks",
"architectures",
"hidden_sizes",
"layers",
"dropouts",
"epochs",
"validation_window",
"holdout_window",
"ensemble_seeds",
"selection_folds",
"learning_rate",
"weight_decay",
"horizon",
"horizons",
"patience",
"context_symbols",
"features",
"seed",
"interval",
"pooled",
"resume_candidate",
}
result = {key: value[key] for key in allowed if key in value}
for key, low, high in (
("limit", 500, 20000),
("epochs", 1, 200),
("validation_window", 64, 2000),
("holdout_window", 64, 1000),
("selection_folds", 1, 12),
("horizon", 1, 96),
("patience", 1, 50),
("seed", 1, 2_147_483_647),
):
if key not in result:
continue
try:
result[key] = max(low, min(high, int(result[key])))
except (TypeError, ValueError):
result.pop(key, None)
if "symbols" in result:
symbols = [
item.strip().upper()
for item in str(result["symbols"]).split(",")
if re.fullmatch(r"[A-Z0-9]{3,20}", item.strip().upper())
]
result["symbols"] = ",".join(symbols[:30])
if "architectures" in result:
architectures = [
item.strip().lower()
for item in str(result["architectures"]).split(",")
if item.strip().lower() in {"lstm", "gru"}
]
result["architectures"] = ",".join(architectures) or "lstm,gru"
for key in (
"lookbacks",
"hidden_sizes",
"layers",
"dropouts",
"ensemble_seeds",
"horizons",
"context_symbols",
"features",
"interval",
):
if key in result:
result[key] = str(result[key])[: 4000 if key == "features" else 500]
for key, low, high in (
("learning_rate", 0.00001, 0.1),
("weight_decay", 0.0, 0.1),
):
if key not in result:
continue
try:
result[key] = max(low, min(high, float(result[key])))
except (TypeError, ValueError):
result.pop(key, None)
if "pooled" in result:
result["pooled"] = result["pooled"] is True
if "resume_candidate" in result:
result["resume_candidate"] = result["resume_candidate"] is True
return result
def _valid_job_id(value: str) -> str:
try:
return str(uuid.UUID(str(value)))
except (ValueError, AttributeError, TypeError) as exc:
raise ValueError("invalid training job id") from exc
def _read_json(path: Path) -> dict[str, Any]:
try:
data = json.loads(path.read_text(encoding="utf-8"))
except (OSError, json.JSONDecodeError) as exc:
raise ValueError(f"invalid training artifact: {path.name}") from exc
if not isinstance(data, dict):
raise ValueError(f"invalid training artifact: {path.name}")
return data
def _validation_passed(validation: dict[str, Any]) -> bool:
if "passed" in validation:
return bool(validation.get("passed"))
return str(validation.get("status", "")).strip().lower() in {"pass", "passed", "ok"}
def _validate_symbol_models(symbols: dict[str, Any]) -> None:
for symbol, entry in symbols.items():
if not isinstance(entry, dict):
raise ValueError(f"candidate model entry is invalid: {symbol}")
if entry.get("model") not in {"torch_lstm", "torch_gru"}:
raise ValueError(f"candidate model architecture is invalid: {symbol}")
try:
lookback = int(entry.get("lookback", 0))
input_size = int(entry.get("input_size", 0))
hidden_size = int(entry.get("hidden_size", 0))
except (TypeError, ValueError) as exc:
raise ValueError(f"candidate model dimensions are invalid: {symbol}") from exc
if not 4 <= lookback <= 512 or not 1 <= input_size <= 256 or not 1 <= hidden_size <= 1024:
raise ValueError(f"candidate model dimensions are out of range: {symbol}")
members = entry.get("ensemble_members")
payloads = members if isinstance(members, list) and members else [entry]
for payload in payloads:
if not isinstance(payload, dict) or not isinstance(payload.get("state_dict"), dict):
raise ValueError(f"candidate recurrent state is missing: {symbol}")
merged = {**entry, **payload}
if merged.get("multitask_head") is True:
required = (
"head_hidden_weight",
"head_hidden_bias",
"return_head_weight",
"return_head_bias",
"event_head_weight",
"event_head_bias",
)
if any(not isinstance(merged.get(name), list) for name in required):
raise ValueError(f"candidate multitask forecast head is missing: {symbol}")
elif not isinstance(merged.get("head_weight"), list) or not isinstance(
merged.get("head_bias"), list
):
raise ValueError(f"candidate forecast head is missing: {symbol}")
def _compact_now() -> str:
return datetime.now(UTC).strftime("%Y%m%dT%H%M%SZ")
def _latest_job(state: dict[str, Any]) -> dict[str, Any] | None:
jobs = state.get("jobs", [])
if not jobs:
return None
latest = jobs[-1]
return latest if isinstance(latest, dict) else None
def _now() -> str:
return datetime.now(UTC).isoformat(timespec="seconds")
def _parse_time(value: str) -> datetime | None:
if not value:
return None
try:
return datetime.fromisoformat(value.replace("Z", "+00:00"))
except ValueError:
return None
def _coerce_percent(value: Any, default: Any = 0) -> int:
try:
number = int(float(value))
except (TypeError, ValueError):
try:
number = int(float(default))
except (TypeError, ValueError):
number = 0
return max(0, min(number, 100))
def _remove_tree(path: Path) -> None:
if not path.exists():
return
for child in path.iterdir():
if child.is_dir():
_remove_tree(child)
else:
child.unlink(missing_ok=True)
path.rmdir()