Skip to main content

wifi_densepose_cli/
room.rs

1//! `enroll` / `train-room` / `room-status` / `room-watch` — ADR-151 Stages 2–5 CLI.
2//!
3//! Drives the `wifi-densepose-calibration` pipeline against a live ESP32 CSI
4//! stream (requires `edge_tier=0` raw CSI). `enroll` walks the guided anchors and
5//! writes labelled features; `train-room` fits the specialist bank; `room-watch`
6//! runs the mixture runtime and prints live room state.
7
8use anyhow::{bail, Result};
9use clap::Args;
10use std::time::{Duration, Instant, SystemTime, UNIX_EPOCH};
11use tokio::net::UdpSocket;
12use wifi_densepose_calibration::{
13    Anchor, AnchorLabel, AnchorQualityGate, AnchorRecorder, EnrollmentEvent, EnrollmentSession,
14    MixtureOfSpecialists, MultiNodeMixture, NodeGeometry, SpecialistBank,
15};
16use wifi_densepose_calibration::extract::{AnchorFeature, Features};
17use wifi_densepose_core::types::CsiFrame;
18use wifi_densepose_signal::BaselineCalibration;
19
20use crate::calibrate::parse_csi_packet;
21
22const RECV_BUF: usize = 2048;
23
24// ---------------------------------------------------------------------------
25// Shared helpers
26// ---------------------------------------------------------------------------
27
28fn now_unix() -> i64 {
29    SystemTime::now()
30        .duration_since(UNIX_EPOCH)
31        .map(|d| d.as_secs() as i64)
32        .unwrap_or(0)
33}
34
35/// Per-frame scalar: mean amplitude across all subcarriers/streams.
36///
37/// Carries presence/motion energy plus the breathing amplitude modulation.
38/// (Validated live on the ESP32 — picks up breathing where a max-variance
39/// subcarrier instead locks onto motion artifacts. A phase-based carrier on a
40/// *stable* subcarrier is the proper higher-SNR refinement — ADR-151 §4.)
41fn frame_scalar(frame: &CsiFrame) -> f32 {
42    let a = &frame.amplitude;
43    if a.is_empty() {
44        return 0.0;
45    }
46    (a.sum() / a.len() as f64) as f32
47}
48
49fn load_baseline(path: &str) -> Result<BaselineCalibration> {
50    let bytes = std::fs::read(path)
51        .map_err(|e| anyhow::anyhow!("cannot read baseline {path}: {e} — run `calibrate` first"))?;
52    BaselineCalibration::from_bytes(&bytes)
53        .map_err(|e| anyhow::anyhow!("invalid baseline {path}: {e}"))
54}
55
56/// Persisted enrollment output (labelled features + audit log).
57#[derive(serde::Serialize, serde::Deserialize)]
58struct EnrollmentData {
59    room_id: String,
60    baseline_id: String,
61    fs_hz: f32,
62    anchors: Vec<AnchorFeature>,
63    session: EnrollmentSession,
64}
65
66// ---------------------------------------------------------------------------
67// enroll
68// ---------------------------------------------------------------------------
69
70/// Arguments for `enroll`.
71#[derive(Args, Debug, Clone)]
72pub struct EnrollArgs {
73    /// UDP port for ESP32 CSI frames (raw CSI; provision with `--edge-tier 0`).
74    #[arg(long, default_value_t = 5005)]
75    pub udp_port: u16,
76    /// Bind address for the UDP socket.
77    #[arg(long, default_value = "0.0.0.0")]
78    pub bind: String,
79    /// Path to the empty-room baseline produced by `calibrate`.
80    #[arg(long, default_value = "./baseline.bin")]
81    pub baseline: String,
82    /// PHY tier (ht20 / ht40 / he20 / he40).
83    #[arg(long, default_value = "ht20")]
84    pub tier: String,
85    /// Room label.
86    #[arg(long, default_value = "default")]
87    pub room_id: String,
88    /// Output enrollment file.
89    #[arg(long, default_value = "./enrollment.json")]
90    pub output: String,
91    /// CSI sample rate (Hz) used for periodicity extraction.
92    #[arg(long, default_value_t = 15.0)]
93    pub fs_hz: f32,
94    /// Max attempts per anchor before moving on.
95    #[arg(long, default_value_t = 2)]
96    pub attempts: u32,
97}
98
99/// Capture one anchor: returns (accepted feature?, anchor verdict, reason).
100async fn capture_anchor(
101    socket: &UdpSocket,
102    baseline: &BaselineCalibration,
103    gate: &AnchorQualityGate,
104    label: AnchorLabel,
105    tier: &str,
106    fs_hz: f32,
107    room_id: &str,
108) -> Result<(Option<AnchorFeature>, Anchor, Option<String>)> {
109    eprintln!("\n[enroll] {} — {}", label.as_str(), label.prompt());
110    for c in (1..=3).rev() {
111        eprintln!("[enroll]   starting in {c}…");
112        tokio::time::sleep(Duration::from_secs(1)).await;
113    }
114    eprintln!("[enroll]   capturing {} s…", label.duration_s());
115
116    let mut recorder = AnchorRecorder::new(label);
117    let mut series: Vec<f32> = Vec::new();
118    let mut buf = vec![0u8; RECV_BUF];
119    let deadline = Instant::now() + Duration::from_secs(label.duration_s() as u64);
120
121    while Instant::now() < deadline {
122        let timeout = Duration::from_millis(500);
123        if let Ok(Ok(n)) = tokio::time::timeout(timeout, socket.recv(&mut buf)).await {
124            if let Some(frame) = parse_csi_packet(&buf[..n], tier) {
125                recorder.record_frame(baseline, &frame);
126                series.push(frame_scalar(&frame));
127            }
128        }
129    }
130
131    let (anchor, reason) = recorder.finalize(gate, now_unix());
132    let feature = if anchor.quality.accepted {
133        Some(AnchorFeature::from_series(room_id, label, &series, fs_hz))
134    } else {
135        None
136    };
137    Ok((feature, anchor, reason))
138}
139
140/// Execute `enroll`.
141pub async fn enroll(args: EnrollArgs) -> Result<()> {
142    let baseline = load_baseline(&args.baseline)?;
143    let baseline_id = baseline.calibration_uuid().to_string();
144    let gate = AnchorQualityGate::default();
145
146    let addr = format!("{}:{}", args.bind, args.udp_port);
147    let socket = UdpSocket::bind(&addr)
148        .await
149        .map_err(|e| anyhow::anyhow!("cannot bind {addr}: {e}"))?;
150    eprintln!("[enroll] room='{}' baseline={} on udp://{addr}", args.room_id, &baseline_id[..8]);
151    eprintln!("[enroll] follow each prompt; bad captures are re-prompted.");
152
153    let mut session = EnrollmentSession::new(&args.room_id, &baseline_id, now_unix());
154    let mut features: Vec<AnchorFeature> = Vec::new();
155
156    for label in AnchorLabel::SEQUENCE {
157        let mut accepted = false;
158        for attempt in 1..=args.attempts {
159            let (feat, anchor, reason) =
160                capture_anchor(&socket, &baseline, &gate, label, &args.tier, args.fs_hz, &args.room_id)
161                    .await?;
162            if anchor.quality.accepted {
163                eprintln!(
164                    "[enroll]   ✓ accepted (presence_z={:.2} motion={:.0}% frames={})",
165                    anchor.quality.presence_z,
166                    anchor.quality.motion_rate * 100.0,
167                    anchor.quality.frames
168                );
169                if let Some(f) = feat {
170                    features.push(f);
171                }
172                session.apply(EnrollmentEvent::AnchorAccepted { anchor });
173                accepted = true;
174                break;
175            } else {
176                let why = reason.unwrap_or_default();
177                eprintln!("[enroll]   ✗ rejected: {why}");
178                session.apply(EnrollmentEvent::AnchorRejected {
179                    label,
180                    reason: why,
181                    at: now_unix(),
182                });
183                if attempt < args.attempts {
184                    eprintln!("[enroll]   retrying ({}/{})…", attempt + 1, args.attempts);
185                }
186            }
187        }
188        if !accepted {
189            eprintln!("[enroll]   moving on without '{}'", label.as_str());
190        }
191    }
192
193    if session.is_complete() {
194        session.apply(EnrollmentEvent::Completed { at: now_unix() });
195    }
196    let (got, total) = session.progress();
197    let data = EnrollmentData {
198        room_id: args.room_id.clone(),
199        baseline_id,
200        fs_hz: args.fs_hz,
201        anchors: features,
202        session,
203    };
204    std::fs::write(
205        &args.output,
206        serde_json::to_string_pretty(&data).map_err(|e| anyhow::anyhow!("serialize: {e}"))?,
207    )
208    .map_err(|e| anyhow::anyhow!("cannot write {}: {e}", args.output))?;
209    eprintln!(
210        "\n[enroll] done: {got}/{total} anchors accepted → {} (next: `train-room`)",
211        args.output
212    );
213    Ok(())
214}
215
216// ---------------------------------------------------------------------------
217// train-room
218// ---------------------------------------------------------------------------
219
220/// Arguments for `train-room`.
221#[derive(Args, Debug, Clone)]
222pub struct TrainRoomArgs {
223    /// Enrollment file from `enroll`.
224    #[arg(long, default_value = "./enrollment.json")]
225    pub enrollment: String,
226    /// Output specialist-bank file.
227    #[arg(long, default_value = "./room-bank.json")]
228    pub output: String,
229    /// Optional transceiver-geometry file: a JSON array of `NodeGeometry`
230    /// records (ADR-152 §2.1.1). Recorded into the enrollment session before
231    /// training so the bank carries the layout it was trained under.
232    #[arg(long)]
233    pub geometry: Option<String>,
234}
235
236/// Execute `train-room`.
237///
238/// If the enrollment session carries a transceiver-geometry snapshot (recorded
239/// at enroll time or supplied here via `--geometry`), it is threaded into the
240/// bank (ADR-152 §2.1.1); a geometry-free enrollment still trains a valid bank.
241pub async fn train_room(args: TrainRoomArgs) -> Result<()> {
242    let raw = std::fs::read_to_string(&args.enrollment)
243        .map_err(|e| anyhow::anyhow!("cannot read {}: {e} — run `enroll` first", args.enrollment))?;
244    let mut data: EnrollmentData =
245        serde_json::from_str(&raw).map_err(|e| anyhow::anyhow!("invalid enrollment: {e}"))?;
246    if data.anchors.is_empty() {
247        bail!("no accepted anchors in {} — re-run enroll", args.enrollment);
248    }
249
250    if let Some(path) = &args.geometry {
251        let graw = std::fs::read_to_string(path)
252            .map_err(|e| anyhow::anyhow!("cannot read geometry {path}: {e}"))?;
253        let geometry: Vec<NodeGeometry> = serde_json::from_str(&graw).map_err(|e| {
254            anyhow::anyhow!("invalid geometry {path}: {e} (expected a JSON array of NodeGeometry records)")
255        })?;
256        data.session.record_geometry(geometry, now_unix());
257    }
258
259    let mut bank = SpecialistBank::train(&data.room_id, &data.baseline_id, &data.anchors, now_unix())
260        .map_err(|e| anyhow::anyhow!("training failed: {e}"))?;
261    match data.session.geometry() {
262        Some(g) if !g.is_empty() => {
263            bank = bank.with_geometry(g.to_vec());
264            eprintln!(
265                "[train-room] geometry: {} node(s) snapshotted into the bank (ADR-152 §2.1.1)",
266                bank.geometry.len()
267            );
268        }
269        _ => eprintln!(
270            "[train-room] no transceiver geometry recorded — bank will not support geometry conditioning (ADR-152 §2.1.2)"
271        ),
272    }
273    std::fs::write(&args.output, bank.to_json().map_err(|e| anyhow::anyhow!("{e}"))?)
274        .map_err(|e| anyhow::anyhow!("cannot write {}: {e}", args.output))?;
275
276    eprintln!(
277        "[train-room] room='{}' trained {} specialists from {} anchors → {}",
278        bank.room_id,
279        bank.trained_kinds().len(),
280        bank.anchor_count,
281        args.output
282    );
283    for k in bank.trained_kinds() {
284        eprintln!("[train-room]   • {k:?}");
285    }
286    Ok(())
287}
288
289// ---------------------------------------------------------------------------
290// room-status
291// ---------------------------------------------------------------------------
292
293/// Arguments for `room-status`.
294#[derive(Args, Debug, Clone)]
295pub struct RoomStatusArgs {
296    /// Specialist-bank file.
297    #[arg(long, default_value = "./room-bank.json")]
298    pub bank: String,
299}
300
301/// Execute `room-status`.
302pub async fn room_status(args: RoomStatusArgs) -> Result<()> {
303    let raw = std::fs::read_to_string(&args.bank)
304        .map_err(|e| anyhow::anyhow!("cannot read {}: {e}", args.bank))?;
305    let bank = SpecialistBank::from_json(&raw).map_err(|e| anyhow::anyhow!("{e}"))?;
306    println!("room:        {}", bank.room_id);
307    println!("baseline:    {}", bank.baseline_id);
308    println!("trained_at:  {}", bank.trained_at_unix_s);
309    println!("anchors:     {}", bank.anchor_count);
310    println!("specialists: {:?}", bank.trained_kinds());
311    Ok(())
312}
313
314// ---------------------------------------------------------------------------
315// room-watch
316// ---------------------------------------------------------------------------
317
318/// Arguments for `room-watch`.
319#[derive(Args, Debug, Clone)]
320pub struct RoomWatchArgs {
321    /// Specialist-bank file (single-node mode).
322    #[arg(long, default_value = "./room-bank.json")]
323    pub bank: String,
324    /// Multistatic mode: map a node id to its bank as `N:path` (repeatable).
325    /// When supplied, frames are grouped by node id and fused (ADR-029/151).
326    #[arg(long = "node-bank", value_name = "N:PATH")]
327    pub node_bank: Vec<String>,
328    /// UDP port for ESP32 CSI frames (raw CSI).
329    #[arg(long, default_value_t = 5005)]
330    pub udp_port: u16,
331    /// Bind address.
332    #[arg(long, default_value = "0.0.0.0")]
333    pub bind: String,
334    /// PHY tier.
335    #[arg(long, default_value = "ht20")]
336    pub tier: String,
337    /// CSI sample rate (Hz).
338    #[arg(long, default_value_t = 15.0)]
339    pub fs_hz: f32,
340    /// Rolling window length (frames) for each inference.
341    #[arg(long, default_value_t = 200)]
342    pub window: usize,
343    /// Seconds to run (0 = until Ctrl-C).
344    #[arg(long, default_value_t = 0)]
345    pub seconds: u32,
346}
347
348/// Execute `room-watch` — live (multistatic) mixture-of-specialists readout.
349pub async fn room_watch(args: RoomWatchArgs) -> Result<()> {
350    if !args.node_bank.is_empty() {
351        return room_watch_multi(args).await;
352    }
353    let raw = std::fs::read_to_string(&args.bank)
354        .map_err(|e| anyhow::anyhow!("cannot read {}: {e}", args.bank))?;
355    let bank = SpecialistBank::from_json(&raw).map_err(|e| anyhow::anyhow!("{e}"))?;
356    let baseline_id = bank.baseline_id.clone();
357    let mix = MixtureOfSpecialists::new(bank);
358
359    let addr = format!("{}:{}", args.bind, args.udp_port);
360    let socket = UdpSocket::bind(&addr)
361        .await
362        .map_err(|e| anyhow::anyhow!("cannot bind {addr}: {e}"))?;
363    eprintln!("[room-watch] inferring on udp://{addr} (window={} frames)", args.window);
364
365    let mut buf = vec![0u8; RECV_BUF];
366    let mut win: std::collections::VecDeque<f32> = std::collections::VecDeque::new();
367    let start = Instant::now();
368    let mut last_print = Instant::now();
369
370    loop {
371        if args.seconds > 0 && start.elapsed() >= Duration::from_secs(args.seconds as u64) {
372            break;
373        }
374        if let Ok(Ok(n)) = tokio::time::timeout(Duration::from_millis(500), socket.recv(&mut buf)).await {
375            if let Some(frame) = parse_csi_packet(&buf[..n], &args.tier) {
376                win.push_back(frame_scalar(&frame));
377                while win.len() > args.window {
378                    win.pop_front();
379                }
380            }
381        }
382        if last_print.elapsed() >= Duration::from_secs(1) && win.len() >= 32 {
383            let series: Vec<f32> = win.iter().copied().collect();
384            let f = Features::from_series(&series, args.fs_hz);
385            let s = mix.infer(&f, &baseline_id);
386            let pres = s.presence.as_ref().map(|r| r.label.clone().unwrap_or_default()).unwrap_or("-".into());
387            let post = s.posture.as_ref().and_then(|r| r.label.clone()).unwrap_or("-".into());
388            let br = s.breathing.as_ref().map(|r| format!("{:.1}bpm", r.value)).unwrap_or("-".into());
389            let hr = s.heartbeat.as_ref().map(|r| format!("{:.0}bpm", r.value)).unwrap_or("-".into());
390            let rest = s.restlessness.as_ref().map(|r| format!("{:.2}", r.value)).unwrap_or("-".into());
391            let flags = format!(
392                "{}{}",
393                if s.vetoed { " VETO" } else { "" },
394                if s.stale { " STALE" } else { "" }
395            );
396            println!(
397                "presence={pres:<7} posture={post:<8} breathing={br:<8} heart={hr:<7} restless={rest}{flags}"
398            );
399            last_print = Instant::now();
400        }
401    }
402    Ok(())
403}
404
405/// Multistatic `room-watch`: fuse several co-located nodes (ADR-029/151).
406async fn room_watch_multi(args: RoomWatchArgs) -> Result<()> {
407    use std::collections::{BTreeMap, VecDeque};
408
409    let mut mix = MultiNodeMixture::new();
410    let mut node_ids: Vec<u8> = Vec::new();
411    for spec in &args.node_bank {
412        let (id_s, path) = spec
413            .split_once(':')
414            .ok_or_else(|| anyhow::anyhow!("--node-bank must be N:path (got {spec:?})"))?;
415        let id: u8 = id_s
416            .parse()
417            .map_err(|_| anyhow::anyhow!("bad node id in {spec:?}"))?;
418        let raw = std::fs::read_to_string(path)
419            .map_err(|e| anyhow::anyhow!("cannot read {path}: {e}"))?;
420        let bank = SpecialistBank::from_json(&raw).map_err(|e| anyhow::anyhow!("{e}"))?;
421        let baseline = bank.baseline_id.clone();
422        mix.add_node(id, bank, baseline);
423        node_ids.push(id);
424    }
425    eprintln!("[room-watch] multistatic over nodes {node_ids:?}");
426
427    let addr = format!("{}:{}", args.bind, args.udp_port);
428    let socket = UdpSocket::bind(&addr)
429        .await
430        .map_err(|e| anyhow::anyhow!("cannot bind {addr}: {e}"))?;
431    eprintln!("[room-watch] fusing on udp://{addr} (window={} frames)", args.window);
432
433    let mut buf = vec![0u8; RECV_BUF];
434    let mut wins: BTreeMap<u8, VecDeque<f32>> = BTreeMap::new();
435    let start = Instant::now();
436    let mut last_print = Instant::now();
437
438    loop {
439        if args.seconds > 0 && start.elapsed() >= Duration::from_secs(args.seconds as u64) {
440            break;
441        }
442        if let Ok(Ok(n)) =
443            tokio::time::timeout(Duration::from_millis(500), socket.recv(&mut buf)).await
444        {
445            if n < 5 {
446                continue;
447            }
448            let node_id = buf[4];
449            if !node_ids.contains(&node_id) {
450                continue;
451            }
452            if let Some(frame) = parse_csi_packet(&buf[..n], &args.tier) {
453                let w = wins.entry(node_id).or_default();
454                w.push_back(frame_scalar(&frame));
455                while w.len() > args.window {
456                    w.pop_front();
457                }
458            }
459        }
460        if last_print.elapsed() >= Duration::from_secs(1) {
461            let per_node: BTreeMap<u8, Features> = wins
462                .iter()
463                .filter(|(_, w)| w.len() >= 32)
464                .map(|(id, w)| {
465                    let series: Vec<f32> = w.iter().copied().collect();
466                    (*id, Features::from_series(&series, args.fs_hz))
467                })
468                .collect();
469            if !per_node.is_empty() {
470                let active: Vec<u8> = per_node.keys().copied().collect();
471                let s = mix.infer(&per_node);
472                let pres = s.presence.as_ref().and_then(|r| r.label.clone()).unwrap_or("-".into());
473                let post = s.posture.as_ref().and_then(|r| r.label.clone()).unwrap_or("-".into());
474                let br = s.breathing.as_ref().map(|r| format!("{:.1}bpm", r.value)).unwrap_or("-".into());
475                let flags = format!(
476                    "{}{}",
477                    if s.vetoed { " VETO" } else { "" },
478                    if s.stale { " STALE" } else { "" }
479                );
480                println!(
481                    "nodes={active:?} presence={pres:<7} posture={post:<8} breathing={br:<8}{flags}"
482                );
483            }
484            last_print = Instant::now();
485        }
486    }
487    Ok(())
488}
489
490#[cfg(test)]
491mod tests {
492    use super::*;
493
494    fn feature(label: AnchorLabel, variance: f32, motion: f32) -> AnchorFeature {
495        AnchorFeature {
496            room_id: "t".into(),
497            label,
498            features: Features {
499                mean: 1.0,
500                variance,
501                motion,
502                breathing_score: 0.0,
503                breathing_hz: 0.0,
504                heart_score: 0.0,
505                heart_hz: 0.0,
506            },
507        }
508    }
509
510    /// Write a minimal valid enrollment file (two anchors, no geometry event).
511    fn write_enrollment(dir: &std::path::Path) -> String {
512        let data = EnrollmentData {
513            room_id: "t".into(),
514            baseline_id: "base-1".into(),
515            fs_hz: 15.0,
516            anchors: vec![
517                feature(AnchorLabel::Empty, 1.0, 0.1),
518                feature(AnchorLabel::StandStill, 10.0, 0.2),
519            ],
520            session: EnrollmentSession::new("t", "base-1", 1000),
521        };
522        let path = dir.join("enrollment.json");
523        std::fs::write(&path, serde_json::to_string(&data).unwrap()).unwrap();
524        path.to_string_lossy().into_owned()
525    }
526
527    fn trained_bank(out: &std::path::Path) -> SpecialistBank {
528        SpecialistBank::from_json(&std::fs::read_to_string(out).unwrap()).unwrap()
529    }
530
531    /// ADR-152 §2.1.1: `--geometry` records into the session and the bank
532    /// snapshots it — enrollment geometry reaches the trained bank.
533    #[tokio::test]
534    async fn train_room_threads_geometry_when_provided() {
535        let dir = tempfile::tempdir().unwrap();
536        let enrollment = write_enrollment(dir.path());
537        let geometry = vec![
538            NodeGeometry::new(1, "tape-measure").with_position(0.0, 0.0, 1.0),
539            NodeGeometry::unknown(2),
540        ];
541        let gpath = dir.path().join("geometry.json");
542        std::fs::write(&gpath, serde_json::to_string(&geometry).unwrap()).unwrap();
543        let out = dir.path().join("bank.json");
544
545        train_room(TrainRoomArgs {
546            enrollment,
547            output: out.to_string_lossy().into_owned(),
548            geometry: Some(gpath.to_string_lossy().into_owned()),
549        })
550        .await
551        .unwrap();
552
553        assert_eq!(trained_bank(&out).geometry, geometry);
554    }
555
556    /// A geometry-free enrollment still trains a valid bank (optional by
557    /// design) — it just carries no snapshot.
558    #[tokio::test]
559    async fn train_room_without_geometry_yields_geometry_free_bank() {
560        let dir = tempfile::tempdir().unwrap();
561        let enrollment = write_enrollment(dir.path());
562        let out = dir.path().join("bank.json");
563
564        train_room(TrainRoomArgs {
565            enrollment,
566            output: out.to_string_lossy().into_owned(),
567            geometry: None,
568        })
569        .await
570        .unwrap();
571
572        let bank = trained_bank(&out);
573        assert!(bank.geometry.is_empty());
574        assert!(bank.presence.is_some(), "bank still trains without geometry");
575    }
576
577    /// Geometry recorded at enroll time (in the session event log) is picked up
578    /// without the `--geometry` flag.
579    #[tokio::test]
580    async fn train_room_uses_session_geometry() {
581        let dir = tempfile::tempdir().unwrap();
582        let geometry = vec![NodeGeometry::new(3, "floor-plan").with_position(1.0, 2.0, 1.5)];
583        let mut session = EnrollmentSession::new("t", "base-1", 1000);
584        session.record_geometry(geometry.clone(), 1000);
585        let data = EnrollmentData {
586            room_id: "t".into(),
587            baseline_id: "base-1".into(),
588            fs_hz: 15.0,
589            anchors: vec![
590                feature(AnchorLabel::Empty, 1.0, 0.1),
591                feature(AnchorLabel::StandStill, 10.0, 0.2),
592            ],
593            session,
594        };
595        let epath = dir.path().join("enrollment.json");
596        std::fs::write(&epath, serde_json::to_string(&data).unwrap()).unwrap();
597        let out = dir.path().join("bank.json");
598
599        train_room(TrainRoomArgs {
600            enrollment: epath.to_string_lossy().into_owned(),
601            output: out.to_string_lossy().into_owned(),
602            geometry: None,
603        })
604        .await
605        .unwrap();
606
607        assert_eq!(trained_bank(&out).geometry, geometry);
608    }
609
610    #[tokio::test]
611    async fn train_room_rejects_invalid_geometry_file() {
612        let dir = tempfile::tempdir().unwrap();
613        let enrollment = write_enrollment(dir.path());
614        let gpath = dir.path().join("geometry.json");
615        std::fs::write(&gpath, r#"{"not":"an array"}"#).unwrap();
616
617        let err = train_room(TrainRoomArgs {
618            enrollment,
619            output: dir.path().join("bank.json").to_string_lossy().into_owned(),
620            geometry: Some(gpath.to_string_lossy().into_owned()),
621        })
622        .await
623        .unwrap_err();
624        assert!(err.to_string().contains("invalid geometry"), "{err}");
625    }
626}