diff --git a/v2/crates/wifi-densepose-sensing-server/src/main.rs b/v2/crates/wifi-densepose-sensing-server/src/main.rs index 8acf6c79..3e8b7a9f 100644 --- a/v2/crates/wifi-densepose-sensing-server/src/main.rs +++ b/v2/crates/wifi-densepose-sensing-server/src/main.rs @@ -20,8 +20,14 @@ mod multistatic_bridge; mod mediatek_csi; mod qualcomm_csi; mod realtek_radar; +mod path_safety; pub mod pose; mod rvf_container; +// ADR-186 (TRAIN-RECONNECT): the in-server training pipeline was written but +// never declared as a module, so it was orphaned / uncompiled. Declaring it +// here compiles it against the real `AppStateInner` and wires its `routes()` +// (including `/ws/train/progress`) into the live router below. +mod training_api; mod rvf_pipeline; mod tracker_bridge; pub mod types; @@ -1120,11 +1126,13 @@ struct AppStateInner { recording_current_id: Option, /// Shutdown signal for the recording writer task. recording_stop_tx: Option>, - // ── Training fields ───────────────────────────────────────────────────── - /// Training status: "idle", "running", "completed", "failed". - training_status: String, - /// Training configuration, if any. - training_config: Option, + // ── Training fields (ADR-186 TRAIN-RECONNECT) ──────────────────────────── + /// Live training state (shared status snapshot + cooperative cancel flag + + /// background task handle) for the in-server trainer in `training_api`. + training_state: training_api::TrainingState, + /// Fan-out channel the background training job publishes progress JSON to; + /// the `/ws/train/progress` WebSocket handler subscribes to it. + training_progress_tx: broadcast::Sender, // ── Adaptive classifier (environment-tuned) ────────────────────────── /// Trained adaptive model (loaded from data/adaptive_model.json or trained at runtime). adaptive_model: Option, @@ -4973,54 +4981,12 @@ fn scan_recording_files() -> Vec { } // ── Training Endpoints ────────────────────────────────────────────────────── - -/// GET /api/v1/train/status — get training status. -async fn train_status(State(state): State) -> Json { - let s = state.read().await; - Json(serde_json::json!({ - "status": s.training_status, - "config": s.training_config, - })) -} - -/// POST /api/v1/train/start — start a training run. -async fn train_start( - State(state): State, - Json(body): Json, -) -> Json { - let mut s = state.write().await; - if s.training_status == "running" { - return Json(serde_json::json!({ - "error": "training already running", - "success": false, - })); - } - s.training_status = "running".to_string(); - s.training_config = Some(body.clone()); - info!("Training started with config: {}", body); - Json(serde_json::json!({ - "success": true, - "status": "running", - "message": "Training pipeline started. Use GET /api/v1/train/status to monitor.", - })) -} - -/// POST /api/v1/train/stop — stop the current training run. -async fn train_stop(State(state): State) -> Json { - let mut s = state.write().await; - if s.training_status != "running" { - return Json(serde_json::json!({ - "error": "no training in progress", - "success": false, - })); - } - s.training_status = "idle".to_string(); - info!("Training stopped"); - Json(serde_json::json!({ - "success": true, - "status": "idle", - })) -} +// +// ADR-186 (TRAIN-RECONNECT): the former stub handlers here flipped a status +// string and logged one line without ever starting a job (issue #1233). They +// are replaced by the real `training_api` router, merged into the app below, +// which runs the pure-Rust trainer on a background task and streams live +// progress over `/ws/train/progress`. // ── Adaptive classifier endpoints ──────────────────────────────────────────── @@ -7822,9 +7788,9 @@ async fn main() { recording_start_time: None, recording_current_id: None, recording_stop_tx: None, - // Training - training_status: "idle".to_string(), - training_config: None, + // Training (ADR-186 TRAIN-RECONNECT) + training_state: training_api::TrainingState::default(), + training_progress_tx: broadcast::channel::(256).0, adaptive_model: adaptive_classifier::AdaptiveModel::load(&adaptive_classifier::model_path()) .ok() @@ -8065,10 +8031,12 @@ async fn main() { .route("/api/v1/recording/start", post(start_recording)) .route("/api/v1/recording/stop", post(stop_recording)) .route("/api/v1/recording/{id}", delete(delete_recording)) - // Training endpoints - .route("/api/v1/train/status", get(train_status)) - .route("/api/v1/train/start", post(train_start)) - .route("/api/v1/train/stop", post(train_stop)) + // Training endpoints (ADR-186 TRAIN-RECONNECT): the real in-server + // trainer + `/ws/train/progress` stream. Merged while the router is + // still `Router` (before `.with_state`) so these routes + // share `AppStateInner` and `/api/v1/train/*` sits under the bearer gate + // applied below (like the rest of `/api/v1/*`). + .merge(training_api::routes()) // Adaptive classifier endpoints .route("/api/v1/adaptive/train", post(adaptive_train)) .route("/api/v1/adaptive/status", get(adaptive_status)) diff --git a/v2/crates/wifi-densepose-sensing-server/src/training_api.rs b/v2/crates/wifi-densepose-sensing-server/src/training_api.rs index fb708a28..c5a21b33 100644 --- a/v2/crates/wifi-densepose-sensing-server/src/training_api.rs +++ b/v2/crates/wifi-densepose-sensing-server/src/training_api.rs @@ -26,7 +26,8 @@ use std::collections::VecDeque; use std::path::PathBuf; -use std::sync::Arc; +use std::sync::atomic::{AtomicBool, Ordering}; +use std::sync::{Arc, Mutex}; use axum::{ extract::{ @@ -38,10 +39,9 @@ use axum::{ Router, }; use serde::{Deserialize, Serialize}; -use tokio::sync::{broadcast, RwLock}; +use tokio::sync::broadcast; use tracing::{error, info, warn}; -use crate::recording::{RecordedFrame, RECORDINGS_DIR}; use crate::rvf_container::RvfBuilder; // ── Constants ──────────────────────────────────────────────────────────────── @@ -49,6 +49,10 @@ use crate::rvf_container::RvfBuilder; /// Directory for trained model output. pub const MODELS_DIR: &str = "data/models"; +/// Directory the training loop reads recorded CSI datasets from. Each +/// `dataset_id` maps to `{RECORDINGS_DIR}/{dataset_id}.csi.jsonl`. +pub const RECORDINGS_DIR: &str = "data/recordings"; + /// Number of COCO keypoints. const N_KEYPOINTS: usize = 17; /// Dimensions per keypoint in the target vector (x, y, z). @@ -67,6 +71,25 @@ const N_GLOBAL_FEATURES: usize = 3; // ── Types ──────────────────────────────────────────────────────────────────── +/// A single recorded CSI frame line, as stored in the `.csi.jsonl` datasets the +/// training loop consumes. +/// +/// This mirrors the on-disk JSONL schema and is intentionally self-contained so +/// the trainer does not couple to the (separate, orphaned) `recording.rs` +/// module. Only the fields the feature extractor needs are read; `rssi` / +/// `noise_floor` / `features` are carried for schema fidelity. +#[derive(Debug, Clone, Serialize, Deserialize)] +pub struct RecordedFrame { + pub timestamp: f64, + pub subcarriers: Vec, + #[serde(default)] + pub rssi: f64, + #[serde(default)] + pub noise_floor: f64, + #[serde(default)] + pub features: serde_json::Value, +} + /// Training configuration submitted with a start request. #[derive(Debug, Clone, Serialize, Deserialize)] pub struct TrainingConfig { @@ -229,24 +252,45 @@ pub struct TrainingProgress { } /// Runtime training state stored in `AppStateInner`. +/// +/// `status` and `cancel` are shared handles (not owned snapshots) so the +/// background training job can update progress and observe stop requests +/// **without holding a reference to the full `AppStateInner`**. That decoupling +/// is what makes the training core ([`run_training_job`]) unit-testable in +/// isolation from the ~60-field server state. pub struct TrainingState { - /// Current status snapshot. - pub status: TrainingStatus, - /// Handle to the background training task (for cancellation). + /// Live status snapshot, shared with the running training job. + pub status: Arc>, + /// Cooperative stop flag; `stop_training` sets it and the job loop observes it. + pub cancel: Arc, + /// Handle to the background training task. pub task_handle: Option>, } impl Default for TrainingState { fn default() -> Self { Self { - status: TrainingStatus::default(), + status: Arc::new(Mutex::new(TrainingStatus::default())), + cancel: Arc::new(AtomicBool::new(false)), task_handle: None, } } } +impl TrainingState { + /// Clone of the current status snapshot. + pub fn snapshot(&self) -> TrainingStatus { + self.status.lock().unwrap().clone() + } + + /// Whether a training job is currently active. + pub fn is_active(&self) -> bool { + self.status.lock().unwrap().active + } +} + /// Shared application state type. -pub type AppState = Arc>; +pub type AppState = Arc>; /// Feature normalization statistics computed from the training set. /// Stored alongside the model weights inside the .rvf container so that @@ -317,11 +361,11 @@ async fn load_recording_frames(dataset_ids: &[String]) -> Vec { all_frames } -/// Attempt to collect frames from the live frame_history buffer in AppState. -/// Each `Vec` in frame_history is a subcarrier amplitude vector. -async fn load_frames_from_history(state: &AppState) -> Vec { - let s = state.read().await; - let history: &VecDeque> = &s.frame_history; +/// Build fallback training frames from a snapshot of the live `frame_history` +/// buffer. Each `Vec` is one frame's subcarrier amplitude vector. Passed as +/// an owned snapshot (not a live `AppState` borrow) so the training core stays +/// state-free and independently testable. +fn frames_from_history(history: &[Vec]) -> Vec { history .iter() .enumerate() @@ -938,13 +982,15 @@ fn deterministic_shuffle(n: usize, seed: u64) -> Vec { /// linear model via mini-batch gradient descent. /// /// On completion, exports a `.rvf` container with real calibrated weights. -async fn real_training_loop( - state: AppState, +async fn run_training_job( + status: Arc>, + cancel: Arc, progress_tx: broadcast::Sender, config: TrainingConfig, dataset_ids: Vec, + history_snapshot: Vec>, training_type: &str, -) { +) -> Option { let total_epochs = config.epochs; let patience = config.early_stopping_patience; let mut best_pck = 0.0f64; @@ -978,7 +1024,7 @@ async fn real_training_loop( let mut frames = load_recording_frames(&dataset_ids).await; if frames.is_empty() { info!("No recordings found for dataset_ids; falling back to live frame_history"); - frames = load_frames_from_history(&state).await; + frames = frames_from_history(&history_snapshot); } if frames.len() < 10 { @@ -999,11 +1045,12 @@ async fn real_training_loop( if let Ok(json) = serde_json::to_string(&fail) { let _ = progress_tx.send(json); } - let mut s = state.write().await; - s.training_state.status.active = false; - s.training_state.status.phase = "failed".to_string(); - s.training_state.task_handle = None; - return; + { + let mut st = status.lock().unwrap(); + st.active = false; + st.phase = "failed".to_string(); + } + return None; } info!("Loaded {} frames for training", frames.len()); @@ -1079,13 +1126,10 @@ async fn real_training_loop( // ── Phase 5: Training loop ─────────────────────────────────────────────── for epoch in 1..=total_epochs { - // Check cancellation. - { - let s = state.read().await; - if !s.training_state.status.active { - info!("Training cancelled at epoch {epoch}"); - break; - } + // Check cancellation (cooperative stop flag set by `stop_training`). + if cancel.load(Ordering::Relaxed) { + info!("Training cancelled at epoch {epoch}"); + break; } let phase = if epoch <= config.warmup_epochs { @@ -1245,10 +1289,10 @@ async fn real_training_loop( let remaining = total_epochs.saturating_sub(epoch); let eta_secs = (remaining as f64 * secs_per_epoch) as u64; - // Update shared state. + // Update the shared status snapshot (read by GET /api/v1/train/status). { - let mut s = state.write().await; - s.training_state.status = TrainingStatus { + let mut st = status.lock().unwrap(); + *st = TrainingStatus { active: true, epoch, total_epochs, @@ -1297,15 +1341,12 @@ async fn real_training_loop( // ── Phase 6: Export .rvf model ─────────────────────────────────────────── - let completed_phase; - { - let s = state.read().await; - completed_phase = if s.training_state.status.active { - "completed" - } else { - "cancelled" - }; - } + let completed_phase = if cancel.load(Ordering::Relaxed) { + "cancelled" + } else { + "completed" + }; + let mut written_rvf: Option = None; // Emit completion message. let completion = TrainingProgress { @@ -1407,28 +1448,32 @@ async fn real_training_loop( }), ); - if let Err(e) = builder.write_to_file(&rvf_path) { - error!("Failed to write trained model RVF: {e}"); - } else { - info!( - "Trained model saved: {} ({} params, pck_torso_h@0.2={:.4})", - rvf_path.display(), - total_params, - best_pck - ); + match builder.write_to_file(&rvf_path) { + Err(e) => { + error!("Failed to write trained model RVF: {e}"); + } + Ok(()) => { + info!( + "Trained model saved: {} ({} params, pck_torso_h@0.2={:.4})", + rvf_path.display(), + total_params, + best_pck + ); + written_rvf = Some(rvf_path); + } } } } - // Mark training as inactive. + // Mark training as inactive in the shared status snapshot. { - let mut s = state.write().await; - s.training_state.status.active = false; - s.training_state.status.phase = completed_phase.to_string(); - s.training_state.task_handle = None; + let mut st = status.lock().unwrap(); + st.active = false; + st.phase = completed_phase.to_string(); } info!("Real {training_type} training finished: phase={completed_phase}"); + written_rvf } // ── Public inference function ──────────────────────────────────────────────── @@ -1565,50 +1610,74 @@ async fn start_training( State(state): State, Json(body): Json, ) -> Json { - // Check if training is already active. - { - let s = state.read().await; - if s.training_state.status.active { - return Json(serde_json::json!({ - "status": "error", - "message": "Training is already active. Stop it first.", - "current_epoch": s.training_state.status.epoch, - "total_epochs": s.training_state.status.total_epochs, - })); - } - } - let config = body.config.clone(); - let dataset_ids = body.dataset_ids.clone(); + match spawn_training_job(&state, config, body.dataset_ids.clone(), "supervised").await { + Ok(()) => Json(serde_json::json!({ + "status": "started", + "type": "supervised", + "dataset_ids": body.dataset_ids, + "config": body.config, + })), + Err(active) => Json(active_error(&active)), + } +} - // Mark training as active and spawn background task. - let progress_tx; - { +/// Snapshot of the already-running job returned when a start is rejected. +fn active_error(snap: &TrainingStatus) -> serde_json::Value { + serde_json::json!({ + "status": "error", + "message": "Training is already active. Stop it first.", + "current_epoch": snap.epoch, + "total_epochs": snap.total_epochs, + }) +} + +/// Seed the shared status, snapshot `frame_history`, and spawn the background +/// training job. Returns `Err(current_status)` if a job is already active. +/// +/// Centralises the single-job guard + spawn used by the supervised, pretrain, +/// and LoRA start handlers so they cannot diverge. +async fn spawn_training_job( + state: &AppState, + config: TrainingConfig, + dataset_ids: Vec, + training_type: &'static str, +) -> Result<(), TrainingStatus> { + let (progress_tx, status, cancel, history_snapshot) = { let s = state.read().await; - progress_tx = s.training_progress_tx.clone(); - } + if s.training_state.is_active() { + return Err(s.training_state.snapshot()); + } + ( + s.training_progress_tx.clone(), + s.training_state.status.clone(), + s.training_state.cancel.clone(), + s.frame_history.iter().cloned().collect::>(), + ) + }; - { - let mut s = state.write().await; - s.training_state.status = TrainingStatus { - active: true, - epoch: 0, - total_epochs: config.epochs, - train_loss: 0.0, - val_pck: 0.0, - val_oks: 0.0, - lr: config.learning_rate, - best_pck: 0.0, - best_epoch: 0, - patience_remaining: config.early_stopping_patience, - eta_secs: None, - phase: "initializing".to_string(), - }; - } + // Clear any prior stop request and seed the initial status snapshot. + cancel.store(false, Ordering::Relaxed); + *status.lock().unwrap() = TrainingStatus { + active: true, + total_epochs: config.epochs, + lr: config.learning_rate, + patience_remaining: config.early_stopping_patience, + phase: "initializing".to_string(), + ..Default::default() + }; - let state_clone = state.clone(); let handle = tokio::spawn(async move { - real_training_loop(state_clone, progress_tx, config, dataset_ids, "supervised").await; + run_training_job( + status, + cancel, + progress_tx, + config, + dataset_ids, + history_snapshot, + training_type, + ) + .await; }); { @@ -1616,57 +1685,46 @@ async fn start_training( s.training_state.task_handle = Some(handle); } - Json(serde_json::json!({ - "status": "started", - "type": "supervised", - "dataset_ids": body.dataset_ids, - "config": body.config, - })) + Ok(()) } async fn stop_training(State(state): State) -> Json { - let mut s = state.write().await; - if !s.training_state.status.active { + let s = state.read().await; + if !s.training_state.is_active() { return Json(serde_json::json!({ "status": "error", "message": "No training is currently active.", })); } - s.training_state.status.active = false; - s.training_state.status.phase = "stopping".to_string(); - - // The background task checks the active flag and will exit. - // We do not abort the handle -- we let it finish the current batch gracefully. + // Set the cooperative stop flag; the background job observes it between + // epochs and exits gracefully after the current batch. We do not abort the + // task handle. + s.training_state.cancel.store(true, Ordering::Relaxed); + { + let mut st = s.training_state.status.lock().unwrap(); + st.phase = "stopping".to_string(); + } + let snap = s.training_state.snapshot(); info!("Training stop requested"); Json(serde_json::json!({ "status": "stopping", - "epoch": s.training_state.status.epoch, - "best_pck": s.training_state.status.best_pck, + "epoch": snap.epoch, + "best_pck": snap.best_pck, })) } async fn training_status(State(state): State) -> Json { let s = state.read().await; - Json(serde_json::to_value(&s.training_state.status).unwrap_or_default()) + Json(serde_json::to_value(s.training_state.snapshot()).unwrap_or_default()) } async fn start_pretrain( State(state): State, Json(body): Json, ) -> Json { - { - let s = state.read().await; - if s.training_state.status.active { - return Json(serde_json::json!({ - "status": "error", - "message": "Training is already active. Stop it first.", - })); - } - } - let config = TrainingConfig { epochs: body.epochs, learning_rate: body.lr, @@ -1675,56 +1733,22 @@ async fn start_pretrain( ..Default::default() }; - let progress_tx; - { - let s = state.read().await; - progress_tx = s.training_progress_tx.clone(); + match spawn_training_job(&state, config, body.dataset_ids.clone(), "pretrain").await { + Ok(()) => Json(serde_json::json!({ + "status": "started", + "type": "pretrain", + "epochs": body.epochs, + "lr": body.lr, + "dataset_ids": body.dataset_ids, + })), + Err(active) => Json(active_error(&active)), } - - { - let mut s = state.write().await; - s.training_state.status = TrainingStatus { - active: true, - total_epochs: body.epochs, - phase: "initializing".to_string(), - ..Default::default() - }; - } - - let state_clone = state.clone(); - let dataset_ids = body.dataset_ids.clone(); - let handle = tokio::spawn(async move { - real_training_loop(state_clone, progress_tx, config, dataset_ids, "pretrain").await; - }); - - { - let mut s = state.write().await; - s.training_state.task_handle = Some(handle); - } - - Json(serde_json::json!({ - "status": "started", - "type": "pretrain", - "epochs": body.epochs, - "lr": body.lr, - "dataset_ids": body.dataset_ids, - })) } async fn start_lora_training( State(state): State, Json(body): Json, ) -> Json { - { - let s = state.read().await; - if s.training_state.status.active { - return Json(serde_json::json!({ - "status": "error", - "message": "Training is already active. Stop it first.", - })); - } - } - let config = TrainingConfig { epochs: body.epochs, learning_rate: 0.0005, // lower LR for LoRA @@ -1735,42 +1759,18 @@ async fn start_lora_training( ..Default::default() }; - let progress_tx; - { - let s = state.read().await; - progress_tx = s.training_progress_tx.clone(); + match spawn_training_job(&state, config, body.dataset_ids.clone(), "lora").await { + Ok(()) => Json(serde_json::json!({ + "status": "started", + "type": "lora", + "base_model_id": body.base_model_id, + "profile_name": body.profile_name, + "rank": body.rank, + "epochs": body.epochs, + "dataset_ids": body.dataset_ids, + })), + Err(active) => Json(active_error(&active)), } - - { - let mut s = state.write().await; - s.training_state.status = TrainingStatus { - active: true, - total_epochs: body.epochs, - phase: "initializing".to_string(), - ..Default::default() - }; - } - - let state_clone = state.clone(); - let dataset_ids = body.dataset_ids.clone(); - let handle = tokio::spawn(async move { - real_training_loop(state_clone, progress_tx, config, dataset_ids, "lora").await; - }); - - { - let mut s = state.write().await; - s.training_state.task_handle = Some(handle); - } - - Json(serde_json::json!({ - "status": "started", - "type": "lora", - "base_model_id": body.base_model_id, - "profile_name": body.profile_name, - "rank": body.rank, - "epochs": body.epochs, - "dataset_ids": body.dataset_ids, - })) } // ── WebSocket handler for training progress ────────────────────────────────── @@ -1792,8 +1792,11 @@ async fn handle_train_ws_client(mut socket: WebSocket, state: AppState) { // Send current status immediately. { - let s = state.read().await; - if let Ok(json) = serde_json::to_string(&s.training_state.status) { + let snapshot = { + let s = state.read().await; + s.training_state.snapshot() + }; + if let Ok(json) = serde_json::to_string(&snapshot) { let msg = serde_json::json!({ "type": "status", "data": serde_json::from_str::(&json).unwrap_or_default(), @@ -2132,4 +2135,142 @@ mod tests { assert_eq!(parsed.n_features, 2); assert_eq!(parsed.mean, vec![1.0, 2.0]); } + + /// Build a small deterministic set of synthetic CSI frames with enough + /// variation that feature extraction is non-degenerate. + fn synthetic_history(n: usize, n_sub: usize) -> Vec> { + (0..n) + .map(|i| { + (0..n_sub) + .map(|k| 10.0 + ((i as f64) * 0.3 + (k as f64) * 0.1).sin() * 2.0) + .collect() + }) + .collect() + } + + /// ADR-186 P3/P6 end-to-end: the real (state-free) training core must + /// (a) stream real progress events over the broadcast channel and + /// (b) actually write a `.rvf` model artifact on completion — not merely + /// flip a status flag. This is the regression guard that keeps the trainer + /// wired (the module was previously orphaned / uncompiled — ADR-186 §1.3). + #[tokio::test] + async fn training_job_streams_real_progress_and_writes_model() { + let history = synthetic_history(40, 56); + + let (tx, mut rx) = broadcast::channel::(1024); + let status = Arc::new(Mutex::new(TrainingStatus::default())); + let cancel = Arc::new(AtomicBool::new(false)); + + let config = TrainingConfig { + epochs: 3, + batch_size: 8, + warmup_epochs: 1, + early_stopping_patience: 10, + ..Default::default() + }; + + // Empty dataset_ids → falls back to the in-memory history snapshot, so + // this test does not depend on the recordings directory. + let rvf = run_training_job( + status.clone(), + cancel, + tx, + config, + Vec::new(), + history, + "supervised", + ) + .await; + + // (b) A real model artifact was produced and exists on disk. + let rvf_path = rvf.expect("training must produce an .rvf model artifact"); + assert!( + rvf_path.exists(), + "rvf artifact should exist at {}", + rvf_path.display() + ); + + // (a) Real progress frames were streamed, at least one carrying an epoch. + let mut n_frames = 0usize; + let mut saw_epoch = false; + let mut saw_completed = false; + while let Ok(msg) = rx.try_recv() { + n_frames += 1; + let v: serde_json::Value = serde_json::from_str(&msg).unwrap(); + if v.get("epoch").and_then(|e| e.as_u64()).unwrap_or(0) >= 1 { + saw_epoch = true; + } + if v.get("phase").and_then(|p| p.as_str()) == Some("completed") { + saw_completed = true; + } + } + assert!(n_frames > 0, "expected streamed progress frames, got none"); + assert!(saw_epoch, "expected at least one epoch-tagged progress frame"); + assert!(saw_completed, "expected a terminal 'completed' progress frame"); + + // Final shared status reflects genuine completion, not just a flag flip: + // real epochs ran (the loop wrote per-epoch status) and a finite loss was + // computed from the real gradient-descent pass. + let final_status = status.lock().unwrap().clone(); + assert!(!final_status.active, "job should be inactive when finished"); + assert_eq!(final_status.phase, "completed"); + assert!( + final_status.epoch >= 1, + "at least one real training epoch should have run" + ); + assert!( + final_status.train_loss.is_finite(), + "a finite training loss should have been computed" + ); + + // Keep the test hermetic — remove the artifact it wrote. + let _ = std::fs::remove_file(&rvf_path); + } + + /// ADR-186 P4 (path safety): a `dataset_id` containing directory traversal + /// is rejected before any file is opened, so the loader returns no frames + /// rather than reading an arbitrary file. + #[tokio::test] + async fn load_recording_frames_rejects_path_traversal() { + let frames = load_recording_frames(&["../../etc/passwd".to_string()]).await; + assert!( + frames.is_empty(), + "path-traversal dataset_id must yield no frames" + ); + } + + /// A job that is cancelled before it starts still exits cleanly and reports + /// the `cancelled` terminal phase (drives `stop_training`'s cooperative flag). + #[tokio::test] + async fn training_job_honors_cancellation() { + let history = synthetic_history(40, 56); + let (tx, _rx) = broadcast::channel::(1024); + let status = Arc::new(Mutex::new(TrainingStatus::default())); + let cancel = Arc::new(AtomicBool::new(true)); // pre-cancelled + + let config = TrainingConfig { + epochs: 50, + batch_size: 8, + warmup_epochs: 1, + early_stopping_patience: 10, + ..Default::default() + }; + + let rvf = run_training_job( + status.clone(), + cancel, + tx, + config, + Vec::new(), + history, + "supervised", + ) + .await; + + // Cancelled before the first epoch → no model, terminal phase cancelled. + assert!(rvf.is_none(), "cancelled run should not export a model"); + let final_status = status.lock().unwrap().clone(); + assert!(!final_status.active); + assert_eq!(final_status.phase, "cancelled"); + } }