1use 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
24fn 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
35fn 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#[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#[derive(Args, Debug, Clone)]
72pub struct EnrollArgs {
73 #[arg(long, default_value_t = 5005)]
75 pub udp_port: u16,
76 #[arg(long, default_value = "0.0.0.0")]
78 pub bind: String,
79 #[arg(long, default_value = "./baseline.bin")]
81 pub baseline: String,
82 #[arg(long, default_value = "ht20")]
84 pub tier: String,
85 #[arg(long, default_value = "default")]
87 pub room_id: String,
88 #[arg(long, default_value = "./enrollment.json")]
90 pub output: String,
91 #[arg(long, default_value_t = 15.0)]
93 pub fs_hz: f32,
94 #[arg(long, default_value_t = 2)]
96 pub attempts: u32,
97}
98
99async 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
140pub 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#[derive(Args, Debug, Clone)]
222pub struct TrainRoomArgs {
223 #[arg(long, default_value = "./enrollment.json")]
225 pub enrollment: String,
226 #[arg(long, default_value = "./room-bank.json")]
228 pub output: String,
229 #[arg(long)]
233 pub geometry: Option<String>,
234}
235
236pub 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#[derive(Args, Debug, Clone)]
295pub struct RoomStatusArgs {
296 #[arg(long, default_value = "./room-bank.json")]
298 pub bank: String,
299}
300
301pub 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#[derive(Args, Debug, Clone)]
320pub struct RoomWatchArgs {
321 #[arg(long, default_value = "./room-bank.json")]
323 pub bank: String,
324 #[arg(long = "node-bank", value_name = "N:PATH")]
327 pub node_bank: Vec<String>,
328 #[arg(long, default_value_t = 5005)]
330 pub udp_port: u16,
331 #[arg(long, default_value = "0.0.0.0")]
333 pub bind: String,
334 #[arg(long, default_value = "ht20")]
336 pub tier: String,
337 #[arg(long, default_value_t = 15.0)]
339 pub fs_hz: f32,
340 #[arg(long, default_value_t = 200)]
342 pub window: usize,
343 #[arg(long, default_value_t = 0)]
345 pub seconds: u32,
346}
347
348pub 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
405async 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 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 #[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 #[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 #[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}