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