Harden trading, training, and monitoring
This commit is contained in:
@@ -4,6 +4,8 @@ import base64
|
||||
import hashlib
|
||||
import json
|
||||
import os
|
||||
import re
|
||||
import shutil
|
||||
import uuid
|
||||
from datetime import UTC
|
||||
from datetime import datetime
|
||||
@@ -20,6 +22,10 @@ ALLOWED_TRAINING_ARTIFACTS = {
|
||||
}
|
||||
RUNNING_TIMEOUT = timedelta(hours=12)
|
||||
ONLINE_WINDOW = timedelta(minutes=3)
|
||||
MAX_ARTIFACT_CHUNK_BYTES = 1024 * 1024
|
||||
MAX_ARTIFACT_BYTES = 64 * 1024 * 1024
|
||||
MAX_ARTIFACT_CHUNKS = 1024
|
||||
REQUIRED_MODEL_BUNDLE = set(ALLOWED_TRAINING_ARTIFACTS)
|
||||
|
||||
|
||||
class TrainingCoordinator:
|
||||
@@ -91,6 +97,7 @@ class TrainingCoordinator:
|
||||
return {"claimed": True, "job": job, "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}")
|
||||
@@ -99,58 +106,87 @@ class TrainingCoordinator:
|
||||
sha256 = str(payload.get("sha256") or "").strip().lower()
|
||||
if index < 0 or total <= 0 or index >= total:
|
||||
raise ValueError("invalid artifact chunk index")
|
||||
if not sha256:
|
||||
raise ValueError("artifact sha256 is required")
|
||||
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")
|
||||
|
||||
chunk_dir = self.upload_root / job_id / name
|
||||
chunk_dir.mkdir(parents=True, exist_ok=True)
|
||||
(chunk_dir / f"{index:06d}.part").write_bytes(chunk)
|
||||
|
||||
if not all((chunk_dir / f"{part:06d}.part").is_file() for part in range(total)):
|
||||
return {"complete": False, "received": index + 1, "total": total}
|
||||
|
||||
target_tmp = self.runtime_dir / f".{name}.{job_id}.tmp"
|
||||
digest = hashlib.sha256()
|
||||
with target_tmp.open("wb") as output:
|
||||
for part in range(total):
|
||||
data = (chunk_dir / f"{part:06d}.part").read_bytes()
|
||||
digest.update(data)
|
||||
output.write(data)
|
||||
if digest.hexdigest().lower() != sha256:
|
||||
target_tmp.unlink(missing_ok=True)
|
||||
raise ValueError("artifact sha256 mismatch")
|
||||
|
||||
self.runtime_dir.mkdir(parents=True, exist_ok=True)
|
||||
os.replace(target_tmp, self.runtime_dir / name)
|
||||
_remove_tree(chunk_dir)
|
||||
|
||||
with self._lock:
|
||||
state = self._load_state()
|
||||
job = self._job_by_id(state, job_id)
|
||||
if job is not None:
|
||||
artifacts = job.setdefault("artifacts", [])
|
||||
artifacts = [item for item in artifacts if item.get("name") != name]
|
||||
artifacts.append({"name": name, "sha256": sha256, "uploaded_at": _now()})
|
||||
job["artifacts"] = artifacts
|
||||
self._save_state(state)
|
||||
return {"complete": True, "name": name, "sha256": sha256}
|
||||
|
||||
def progress(self, job_id: str, payload: dict[str, Any] | None = None) -> dict[str, Any]:
|
||||
payload = payload or {}
|
||||
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")
|
||||
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
|
||||
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
|
||||
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")
|
||||
if isinstance(payload.get("worker"), dict):
|
||||
state["worker"] = self._worker_from_payload(payload["worker"])
|
||||
job["status"] = str(payload.get("status") or job.get("status") or "running")
|
||||
job["phase"] = str(payload.get("phase") or job.get("phase") or "running")
|
||||
job["message"] = str(payload.get("message") or job.get("message") or "")
|
||||
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):
|
||||
@@ -160,12 +196,18 @@ class TrainingCoordinator:
|
||||
|
||||
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")
|
||||
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))
|
||||
@@ -176,6 +218,65 @@ class TrainingCoordinator:
|
||||
self._save_state(state)
|
||||
return {"ok": True, "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"))
|
||||
@@ -262,8 +363,97 @@ class TrainingCoordinator:
|
||||
def _safe_parameters(value: Any) -> dict[str, Any]:
|
||||
if not isinstance(value, dict):
|
||||
return {}
|
||||
allowed = {"symbols", "limit", "lookbacks", "architectures", "hidden_sizes", "layers", "dropouts", "epochs"}
|
||||
return {key: value[key] for key in allowed if key in value}
|
||||
allowed = {
|
||||
"symbols",
|
||||
"limit",
|
||||
"lookbacks",
|
||||
"architectures",
|
||||
"hidden_sizes",
|
||||
"layers",
|
||||
"dropouts",
|
||||
"epochs",
|
||||
"holdout_window",
|
||||
"resume_candidate",
|
||||
}
|
||||
result = {key: value[key] for key in allowed if key in value}
|
||||
for key, low, high in (
|
||||
("limit", 500, 5000),
|
||||
("epochs", 1, 200),
|
||||
("holdout_window", 64, 1000),
|
||||
):
|
||||
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"):
|
||||
if key in result:
|
||||
result[key] = str(result[key])[:200]
|
||||
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}")
|
||||
if not isinstance(entry.get("state_dict"), dict):
|
||||
raise ValueError(f"candidate recurrent state is missing: {symbol}")
|
||||
if not isinstance(entry.get("head_weight"), list) or not isinstance(entry.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:
|
||||
|
||||
Reference in New Issue
Block a user