From 21989e36400cf425ad49bb51eecfe9cdc793ec76 Mon Sep 17 00:00:00 2001 From: chopratejas Date: Fri, 1 May 2026 10:06:56 -0700 Subject: [PATCH] fix(rust): port MessageScorer to Rust + parity harness (PR-A) MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit Direct port of `headroom.transforms.scoring.MessageScorer` (459 LOC). Foundation piece for the IntelligentContext port (PR-B onward). What's wired: - Deterministic factors fully ported: recency (exp-decay), forward references (tool_call_id graph), token density (unique/total). - External-dep factors gated behind traits: `EmbeddingProvider` and `ToinProvider`. No concrete impls yet — both default to neutral values matching Python's `embedding_provider=None` / `toin=None`. PR-A1 wires fastembed; PR-A2 plugs in a PyO3 ToinProvider. - ScoringWeights + MessageScore with serde + BTreeMap-ordered breakdown for stable JSON. Parity: - 13 fixtures recorded from Python, byte-equal under the comparator. - Floats rounded to 5 decimals on both sides — absorbs f32-vs-f64 drift in the weighted sum without masking real bugs. Drive-by: re-fix three pre-existing clippy errors in smart_crusher/crusher.rs that re-emerged with new test additions (field_reassign_with_default + dead hash_array_for_ccr). --- crates/headroom-core/src/lib.rs | 1 + crates/headroom-core/src/scoring/mod.rs | 48 ++ crates/headroom-core/src/scoring/score.rs | 123 +++ crates/headroom-core/src/scoring/scorer.rs | 736 ++++++++++++++++++ crates/headroom-core/src/scoring/traits.rs | 175 +++++ crates/headroom-core/src/scoring/weights.rs | 170 ++++ .../src/transforms/smart_crusher/crusher.rs | 27 +- crates/headroom-parity/src/lib.rs | 125 +++ ...om_weights_recency_heavy_5cc105616d48.json | 125 +++ .../drop_safe_truth_table_2e03ebdc1b01.json | 133 ++++ .../message_scorer/empty_c94d65311b40.json | 16 + .../high_density_dd454e2933e3.json | 43 + .../linear_5_messages_71f1ec6a317c.json | 145 ++++ .../low_density_28ad2371fdc0.json | 43 + .../multi_ref_to_assistant_67a047a411c4.json | 171 ++++ .../non_string_content_9ab4060e7b90.json | 76 ++ .../orphan_tool_response_adb4e13b7541.json | 69 ++ .../single_user_message_2b4fac2277be.json | 43 + .../slow_decay_de4fafb3a42b.json | 218 ++++++ .../tool_call_pair_3443697600a6.json | 134 ++++ .../unicode_content_22231f9bffa6.json | 68 ++ tests/parity/record_message_scorer.py | 350 +++++++++ 22 files changed, 3024 insertions(+), 15 deletions(-) create mode 100644 crates/headroom-core/src/scoring/mod.rs create mode 100644 crates/headroom-core/src/scoring/score.rs create mode 100644 crates/headroom-core/src/scoring/scorer.rs create mode 100644 crates/headroom-core/src/scoring/traits.rs create mode 100644 crates/headroom-core/src/scoring/weights.rs create mode 100644 tests/parity/fixtures/message_scorer/custom_weights_recency_heavy_5cc105616d48.json create mode 100644 tests/parity/fixtures/message_scorer/drop_safe_truth_table_2e03ebdc1b01.json create mode 100644 tests/parity/fixtures/message_scorer/empty_c94d65311b40.json create mode 100644 tests/parity/fixtures/message_scorer/high_density_dd454e2933e3.json create mode 100644 tests/parity/fixtures/message_scorer/linear_5_messages_71f1ec6a317c.json create mode 100644 tests/parity/fixtures/message_scorer/low_density_28ad2371fdc0.json create mode 100644 tests/parity/fixtures/message_scorer/multi_ref_to_assistant_67a047a411c4.json create mode 100644 tests/parity/fixtures/message_scorer/non_string_content_9ab4060e7b90.json create mode 100644 tests/parity/fixtures/message_scorer/orphan_tool_response_adb4e13b7541.json create mode 100644 tests/parity/fixtures/message_scorer/single_user_message_2b4fac2277be.json create mode 100644 tests/parity/fixtures/message_scorer/slow_decay_de4fafb3a42b.json create mode 100644 tests/parity/fixtures/message_scorer/tool_call_pair_3443697600a6.json create mode 100644 tests/parity/fixtures/message_scorer/unicode_content_22231f9bffa6.json create mode 100644 tests/parity/record_message_scorer.py diff --git a/crates/headroom-core/src/lib.rs b/crates/headroom-core/src/lib.rs index 5090f7d65..a80534abe 100644 --- a/crates/headroom-core/src/lib.rs +++ b/crates/headroom-core/src/lib.rs @@ -2,6 +2,7 @@ pub mod ccr; pub mod relevance; +pub mod scoring; pub mod signals; pub mod tokenizer; pub mod transforms; diff --git a/crates/headroom-core/src/scoring/mod.rs b/crates/headroom-core/src/scoring/mod.rs new file mode 100644 index 000000000..dc213f8a4 --- /dev/null +++ b/crates/headroom-core/src/scoring/mod.rs @@ -0,0 +1,48 @@ +//! Message-level importance scoring — used by IntelligentContextManager. +//! +//! # Why this lives at the crate root (parallel to `signals/`) +//! +//! `signals/` scores **lines** for line-level compressors (logs, search, +//! diffs). `scoring/` scores **messages** for conversation-level context +//! management. Different inputs, different consumers — separating them +//! keeps each trait surface clean and lets future ports compose without +//! a giant `signals` god-module. +//! +//! # Port status (Phase 7g PR-A, 2026-04-30) +//! +//! Direct port of `headroom/transforms/scoring.py` (459 LOC). The +//! deterministic factors — recency, forward references, density — +//! are fully implemented and parity-tested against Python. The +//! external-dependency factors are gated behind trait surfaces: +//! +//! - **TOIN** (`ToinProvider` trait): no concrete impl yet. Calls +//! return `0.5` for `toin_importance` and `0.0` for `error_indicator` +//! when no provider is wired in. PR-A1 will plug in a `PyO3` +//! `ToinProvider` so Rust can read Python's TOIN state. +//! - **Embeddings** (`EmbeddingProvider` trait): no concrete impl +//! here yet (the crate already has `relevance::EmbeddingScorer` +//! for SmartCrusher; PR-A1 wires the same `bge-small-en-v1.5` +//! model into a `MessageEmbedder` adapter). Until then, +//! `semantic_score` returns `0.5` (neutral). +//! +//! The trait surface is FULL — when the providers land, no API +//! changes are needed inside `MessageScorer`. The neutral-value +//! defaults match Python behavior when those subsystems are not +//! configured (`toin=None`, `embedding_provider=None`). +//! +//! # No hardcoded patterns (project convention) +//! +//! Mirrors the Python module's design principle: importance derives +//! from computed metrics (recency/density/refs), TOIN-learned +//! patterns (field semantics, retrieval rates), and embedding +//! similarity. No keyword regex, no hardcoded "error" strings. + +pub mod score; +pub mod scorer; +pub mod traits; +pub mod weights; + +pub use score::MessageScore; +pub use scorer::MessageScorer; +pub use traits::{EmbeddingProvider, ToinFieldSemantic, ToinPattern, ToinProvider}; +pub use weights::ScoringWeights; diff --git a/crates/headroom-core/src/scoring/score.rs b/crates/headroom-core/src/scoring/score.rs new file mode 100644 index 000000000..11446e4a7 --- /dev/null +++ b/crates/headroom-core/src/scoring/score.rs @@ -0,0 +1,123 @@ +//! `MessageScore` — scoring output for a single message. +//! +//! Mirrors `headroom.transforms.scoring.MessageScore` byte-for-byte +//! including the per-component breakdown. The breakdown is what +//! IntelligentContextManager logs for debug + what TOIN consumes +//! for learning. + +use std::collections::BTreeMap; + +use serde::{Deserialize, Serialize}; + +/// Importance score for a single message. +/// +/// All component scores are in the range `[0.0, 1.0]` where higher = +/// more important. `total_score` is a weighted sum of the components +/// using [`crate::scoring::ScoringWeights`]. +/// +/// # Determinism +/// +/// For the deterministic factors (recency, forward_reference, +/// token_density), the score is a pure function of the input +/// messages + index. Two runs with the same input produce +/// byte-identical scores. +/// +/// For the external-dep factors (semantic_score, toin_score, +/// error_score), the value depends on whether providers are wired +/// in. Without providers, these return neutral defaults (`0.5` / +/// `0.0`) — same as Python's behavior with `embedding_provider=None` +/// / `toin=None`. +#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)] +pub struct MessageScore { + pub message_index: usize, + pub total_score: f32, + + pub recency_score: f32, + pub semantic_score: f32, + pub toin_score: f32, + pub error_score: f32, + pub reference_score: f32, + pub density_score: f32, + + /// Estimated tokens for this message. Python uses + /// `len(content) // 4` as a rough heuristic; we mirror exactly. + /// For non-string content (e.g. tool_calls list), Python returns + /// `100` as a default — we mirror that too. + pub tokens: usize, + + pub is_protected: bool, + /// `not in_tool_unit OR not protected`. Mirrors Python's + /// confusing-but-faithful definition (it's intentionally the OR + /// of two negatives — see Python's `MessageScorer._score_message` + /// line 189). + pub drop_safe: bool, + + /// Per-factor breakdown for debug logging + TOIN learning. + /// Keyed `BTreeMap` for deterministic JSON serialization order + /// (Python's `dict` preserves insertion order; we want stable + /// alphabetical so parity-fixture diffs are stable across runs). + pub score_breakdown: BTreeMap, +} + +#[cfg(test)] +mod tests { + use super::*; + + #[test] + fn round_trips_through_serde() { + let mut breakdown = BTreeMap::new(); + breakdown.insert("recency".to_string(), 0.9); + breakdown.insert("semantic".to_string(), 0.5); + breakdown.insert("toin".to_string(), 0.5); + breakdown.insert("error".to_string(), 0.0); + breakdown.insert("reference".to_string(), 0.0); + breakdown.insert("density".to_string(), 0.7); + + let s = MessageScore { + message_index: 3, + total_score: 0.42, + recency_score: 0.9, + semantic_score: 0.5, + toin_score: 0.5, + error_score: 0.0, + reference_score: 0.0, + density_score: 0.7, + tokens: 25, + is_protected: false, + drop_safe: true, + score_breakdown: breakdown, + }; + + let json = serde_json::to_string(&s).unwrap(); + let back: MessageScore = serde_json::from_str(&json).unwrap(); + assert_eq!(s, back); + } + + #[test] + fn breakdown_serializes_in_alphabetical_order() { + // The breakdown order matters for parity-fixture stability. + // BTreeMap iteration is alphabetical → parity diffs stay + // deterministic across runs. + let mut breakdown = BTreeMap::new(); + // Insert in non-alphabetical order: + for (k, v) in [ + ("toin", 1.0), + ("recency", 2.0), + ("density", 3.0), + ("semantic", 4.0), + ("error", 5.0), + ("reference", 6.0), + ] { + breakdown.insert(k.to_string(), v); + } + let json = serde_json::to_string(&breakdown).unwrap(); + // density < error < recency < reference < semantic < toin alphabetically + assert!( + json.find("density").unwrap() < json.find("error").unwrap() + && json.find("error").unwrap() < json.find("recency").unwrap() + && json.find("recency").unwrap() < json.find("reference").unwrap() + && json.find("reference").unwrap() < json.find("semantic").unwrap() + && json.find("semantic").unwrap() < json.find("toin").unwrap() + ); + } +} diff --git a/crates/headroom-core/src/scoring/scorer.rs b/crates/headroom-core/src/scoring/scorer.rs new file mode 100644 index 000000000..94b3ce8dc --- /dev/null +++ b/crates/headroom-core/src/scoring/scorer.rs @@ -0,0 +1,736 @@ +//! `MessageScorer` — six-factor importance scoring for messages. +//! +//! Direct port of `headroom.transforms.scoring.MessageScorer`. See +//! the module-level doc in `mod.rs` for what is and isn't wired up +//! in PR-A. +//! +//! # Parity contract +//! +//! For the deterministic factors (recency, forward_reference, +//! token_density), the Rust implementation must produce +//! float-epsilon-equal scores to the Python implementation given +//! the same input. This is what `crates/headroom-parity/` validates. +//! +//! For non-deterministic factors (semantic, toin, error), parity is +//! validated only when both implementations are configured with the +//! same providers. Without providers, both return the same neutral +//! defaults. + +use std::collections::HashMap; +use std::sync::Mutex; + +use serde_json::{Map, Value}; + +use crate::scoring::score::MessageScore; +use crate::scoring::traits::{EmbeddingProvider, ToinPattern, ToinProvider}; +use crate::scoring::weights::ScoringWeights; + +/// Six-factor importance scorer. See module-level docs for the +/// list of factors and their weights. +/// +/// # Thread safety +/// +/// `MessageScorer` is `Send + Sync`. The internal embedding cache is +/// guarded by a `Mutex` — contention is low because cache writes +/// happen at most once per (message, scorer) pair and scorer +/// instances are typically per-request, not shared. +pub struct MessageScorer { + weights: ScoringWeights, + toin: Option>, + embedding_provider: Option>, + recency_decay_rate: f32, + embedding_cache: Mutex>>, +} + +impl MessageScorer { + /// Create a new scorer. + /// + /// `weights` are normalized on construction (matching Python's + /// `ScoringWeights().normalized()` in `__init__`). Pass `None` + /// for `toin` and `embedding_provider` to use neutral defaults + /// for those factors — same as Python's `toin=None` / + /// `embedding_provider=None`. + pub fn new( + weights: Option, + toin: Option>, + embedding_provider: Option>, + recency_decay_rate: f32, + ) -> Self { + Self { + weights: weights.unwrap_or_default().normalized(), + toin, + embedding_provider, + recency_decay_rate, + embedding_cache: Mutex::new(HashMap::new()), + } + } + + /// Create a scorer with default weights, no providers, and the + /// Python-default decay rate of 0.1. Useful for tests and the + /// deterministic-only code path. + pub fn with_defaults() -> Self { + Self::new(None, None, None, 0.1) + } + + /// Score every message in the list. + /// + /// `protected_indices` are indices the caller has marked as + /// system / pinned / never-drop. They get `is_protected=true`. + /// `tool_unit_indices` are part of an inseparable tool unit + /// (assistant tool_call + tool response pair) — these get + /// `drop_safe=false` unless also protected, matching Python's + /// confusing-but-faithful `not in_tool_unit OR not protected` + /// rule. + pub fn score_messages( + &self, + messages: &[Value], + protected_indices: &std::collections::HashSet, + tool_unit_indices: &std::collections::HashSet, + ) -> Vec { + let forward_refs = Self::compute_forward_references(messages); + let recent_embedding = self.compute_recent_context_embedding(messages, 3); + + messages + .iter() + .enumerate() + .map(|(i, msg)| { + self.score_message( + msg, + i, + messages.len(), + protected_indices.contains(&i), + tool_unit_indices.contains(&i), + &forward_refs, + recent_embedding.as_deref(), + ) + }) + .collect() + } + + #[allow(clippy::too_many_arguments)] + fn score_message( + &self, + msg: &Value, + index: usize, + total: usize, + protected: bool, + in_tool_unit: bool, + forward_refs: &HashMap, + recent_embedding: Option<&[f32]>, + ) -> MessageScore { + let recency = self.compute_recency_score(index, total); + let semantic = self.compute_semantic_score(msg, index, recent_embedding); + let toin = self.compute_toin_score(msg); + let error = self.compute_error_score(msg); + let reference = self.compute_reference_score(index, forward_refs); + let density = Self::compute_density_score(msg); + + let w = &self.weights; + let total_score = w.recency * recency + + w.semantic_similarity * semantic + + w.toin_importance * toin + + w.error_indicator * error + + w.forward_reference * reference + + w.token_density * density; + + let tokens = estimate_tokens(msg); + + // BTreeMap so JSON serialization order is alphabetical. + let mut breakdown = std::collections::BTreeMap::new(); + breakdown.insert("recency".to_string(), recency); + breakdown.insert("semantic".to_string(), semantic); + breakdown.insert("toin".to_string(), toin); + breakdown.insert("error".to_string(), error); + breakdown.insert("reference".to_string(), reference); + breakdown.insert("density".to_string(), density); + + MessageScore { + message_index: index, + total_score, + recency_score: recency, + semantic_score: semantic, + toin_score: toin, + error_score: error, + reference_score: reference, + density_score: density, + tokens, + is_protected: protected, + // Python: `not in_tool_unit or not protected` — see + // scoring.py:189. This *is* the literal expression. + drop_safe: !in_tool_unit || !protected, + score_breakdown: breakdown, + } + } + + fn compute_recency_score(&self, index: usize, total: usize) -> f32 { + if total <= 1 { + return 1.0; + } + let position_from_end = (total - 1 - index) as f32; + (-self.recency_decay_rate * position_from_end).exp() + } + + fn compute_semantic_score( + &self, + msg: &Value, + index: usize, + recent_embedding: Option<&[f32]>, + ) -> f32 { + let Some(provider) = self.embedding_provider.as_ref() else { + return 0.5; + }; + let Some(recent) = recent_embedding else { + return 0.5; + }; + + let content = match msg.get("content") { + Some(Value::String(s)) if !s.trim().is_empty() => s, + _ => return 0.5, + }; + + // Cache lookup; populate on miss. + let msg_embedding: Vec = { + let mut cache = self.embedding_cache.lock().unwrap(); + if let Some(cached) = cache.get(&index) { + cached.clone() + } else { + match provider.embed(content) { + Ok(v) => { + cache.insert(index, v.clone()); + v + } + Err(_) => return 0.5, + } + } + }; + + cosine_similarity(&msg_embedding, recent) + } + + fn compute_toin_score(&self, msg: &Value) -> f32 { + let Some(toin) = self.toin.as_ref() else { + return 0.5; + }; + if msg.get("role").and_then(Value::as_str) != Some("tool") { + return 0.5; + } + let Some(content_value) = parse_tool_content(msg) else { + return 0.5; + }; + + let Some(pattern) = toin.pattern_for_tool_content(&content_value) else { + return 0.5; + }; + + if pattern.confidence < 0.3 { + return 0.5; + } + + let mut score = 0.5 + pattern.retrieval_rate * 0.5; + + if !pattern.commonly_retrieved_fields.is_empty() { + // Python: min(0.1, 0.02 * len(commonly_retrieved_fields)) + let boost = (0.02 * pattern.commonly_retrieved_fields.len() as f32).min(0.1); + score = (score + boost).min(1.0); + } + + score + } + + fn compute_error_score(&self, msg: &Value) -> f32 { + let Some(toin) = self.toin.as_ref() else { + return 0.0; + }; + if msg.get("role").and_then(Value::as_str) != Some("tool") { + return 0.0; + } + let Some(content_value) = parse_tool_content(msg) else { + return 0.0; + }; + let Some(pattern) = toin.pattern_for_tool_content(&content_value) else { + return 0.0; + }; + + let (error_field_count, high_confidence_errors) = count_error_fields(&pattern); + + if error_field_count == 0 { + return 0.0; + } + + // Python: base_score = min(1.0, 0.3 * error_field_count) + // confidence_boost = min(0.5, 0.2 * high_confidence_errors) + let base = (0.3 * error_field_count as f32).min(1.0); + let boost = (0.2 * high_confidence_errors as f32).min(0.5); + base + boost + } + + fn compute_reference_score(&self, index: usize, forward_refs: &HashMap) -> f32 { + let count = *forward_refs.get(&index).unwrap_or(&0); + if count == 0 { + return 0.0; + } + // Python: min(1.0, 0.3 + 0.2 * math.log(ref_count + 1)) + // math.log is natural log. + let v = 0.3 + 0.2 * ((count + 1) as f32).ln(); + v.min(1.0) + } + + fn compute_density_score(msg: &Value) -> f32 { + let content = match msg.get("content") { + Some(Value::String(s)) => s, + _ => return 0.5, + }; + // Python: `len(content) < 10` — len counts chars (code points). + if content.chars().count() < 10 { + return 0.5; + } + let lower = content.to_lowercase(); + let tokens: Vec<&str> = lower.split_whitespace().collect(); + if tokens.len() < 3 { + return 0.5; + } + let unique: std::collections::HashSet<&&str> = tokens.iter().collect(); + let density = unique.len() as f32 / tokens.len() as f32; + // Python: min(1.0, max(0.0, (density - 0.2) / 0.6)) + ((density - 0.2) / 0.6).clamp(0.0, 1.0) + } + + fn compute_forward_references(messages: &[Value]) -> HashMap { + let mut refs: HashMap = HashMap::new(); + let mut tool_call_ids: HashMap = HashMap::new(); + + for (i, msg) in messages.iter().enumerate() { + let role = msg.get("role").and_then(Value::as_str); + match role { + Some("assistant") => { + if let Some(tcs) = msg.get("tool_calls").and_then(Value::as_array) { + for tc in tcs { + if let Some(id) = tc.get("id").and_then(Value::as_str) { + tool_call_ids.insert(id.to_string(), i); + } + } + } + } + Some("tool") => { + if let Some(tcid) = msg.get("tool_call_id").and_then(Value::as_str) { + if let Some(&ref_idx) = tool_call_ids.get(tcid) { + *refs.entry(ref_idx).or_insert(0) += 1; + } + } + } + _ => {} + } + } + refs + } + + fn compute_recent_context_embedding( + &self, + messages: &[Value], + num_recent: usize, + ) -> Option> { + let provider = self.embedding_provider.as_ref()?; + let start = messages.len().saturating_sub(num_recent); + let mut texts: Vec<&str> = Vec::new(); + for msg in &messages[start..] { + if let Some(Value::String(s)) = msg.get("content") { + if !s.trim().is_empty() { + texts.push(s); + } + } + } + if texts.is_empty() { + return None; + } + let combined = texts.join(" "); + provider.embed(&combined).ok() + } +} + +/// Token estimate. Python: `len(content) // 4` for strings, `100` +/// for non-strings. We use `chars().count()` (code points) to match +/// Python's `len(str)`. +fn estimate_tokens(msg: &Value) -> usize { + match msg.get("content") { + Some(Value::String(s)) => s.chars().count() / 4, + _ => 100, + } +} + +/// Parse a tool message's `content` as JSON if it's a string, or +/// return it directly if already an object/array. Matches Python's +/// behavior in `_compute_toin_score` / `_compute_error_score`: +/// returns `None` for anything that isn't a list/dict. +fn parse_tool_content(msg: &Value) -> Option { + let content = msg.get("content")?; + let parsed = match content { + Value::String(s) => { + if s.is_empty() { + return None; + } + serde_json::from_str::(s).ok()? + } + v => v.clone(), + }; + match &parsed { + Value::Object(_) | Value::Array(_) => { + // Python: `if not items: return 0.5/0.0` — empty + // list/object disqualifies. + if let Value::Array(a) = &parsed { + if a.is_empty() { + return None; + } + } + if let Value::Object(o) = &parsed { + if o.is_empty() { + // Python wraps a dict into a list before checking + // truthiness; an empty dict becomes `[{}]` which + // is truthy. Mirror exactly: + let _ = o; + // (still return Some since [{}] is truthy in Py) + } + } + Some(parsed) + } + _ => None, + } +} + +/// Count `error_indicator` fields in a TOIN pattern. Returns +/// `(total_error_fields, high_confidence_error_fields)` where +/// high-confidence means `confidence >= 0.7`. +fn count_error_fields(pattern: &ToinPattern) -> (u32, u32) { + let mut total = 0u32; + let mut high = 0u32; + for field_sem in pattern.field_semantics.values() { + if field_sem.inferred_type == "error_indicator" { + total += 1; + if field_sem.confidence >= 0.7 { + high += 1; + } + } + } + (total, high) +} + +fn cosine_similarity(a: &[f32], b: &[f32]) -> f32 { + if a.len() != b.len() || a.is_empty() { + return 0.0; + } + let mut dot = 0.0f32; + let mut na = 0.0f32; + let mut nb = 0.0f32; + for i in 0..a.len() { + dot += a[i] * b[i]; + na += a[i] * a[i]; + nb += b[i] * b[i]; + } + if na == 0.0 || nb == 0.0 { + return 0.0; + } + dot / (na.sqrt() * nb.sqrt()) +} + +/// Helper to build a message `Value` for tests. Public-in-crate so +/// integration tests can use it too. +#[allow(dead_code)] +pub(crate) fn msg(role: &str, content: &str) -> Value { + let mut m = Map::new(); + m.insert("role".to_string(), Value::String(role.to_string())); + m.insert("content".to_string(), Value::String(content.to_string())); + Value::Object(m) +} + +#[cfg(test)] +mod tests { + use super::*; + use std::collections::HashSet; + + fn empty_set() -> HashSet { + HashSet::new() + } + + #[test] + fn recency_single_message_returns_one() { + let s = MessageScorer::with_defaults(); + assert_eq!(s.compute_recency_score(0, 1), 1.0); + assert_eq!(s.compute_recency_score(0, 0), 1.0); + } + + #[test] + fn recency_last_message_full_score() { + let s = MessageScorer::with_defaults(); + // Last message: position_from_end = 0, score = e^0 = 1.0 + assert!((s.compute_recency_score(4, 5) - 1.0).abs() < 1e-6); + } + + #[test] + fn recency_decays_exponentially() { + let s = MessageScorer::with_defaults(); + // total=10, index=0 → position_from_end=9 → e^(-0.1*9) = e^-0.9 + let expected = (-0.9f32).exp(); + let got = s.compute_recency_score(0, 10); + assert!( + (got - expected).abs() < 1e-6, + "expected {expected}, got {got}" + ); + } + + #[test] + fn density_short_content_is_neutral() { + // Less than 10 chars + let m = msg("user", "hi"); + assert_eq!(MessageScorer::compute_density_score(&m), 0.5); + } + + #[test] + fn density_few_tokens_is_neutral() { + // >= 10 chars but < 3 tokens after split + let m = msg("user", "abcdefghij"); // single token, 10 chars + assert_eq!(MessageScorer::compute_density_score(&m), 0.5); + } + + #[test] + fn density_all_unique_clamps_to_one() { + // 10 unique tokens, 10 total → density=1.0 → (1.0-0.2)/0.6=1.33 → clamped to 1.0 + let m = msg( + "user", + "alpha bravo charlie delta echo foxtrot golf hotel india juliet", + ); + assert!((MessageScorer::compute_density_score(&m) - 1.0).abs() < 1e-6); + } + + #[test] + fn density_repeated_tokens_lowers_score() { + // 4 unique / 8 total = 0.5 → (0.5-0.2)/0.6 = 0.5 + let m = msg( + "user", + "alpha bravo charlie delta alpha bravo charlie delta", + ); + let got = MessageScorer::compute_density_score(&m); + assert!((got - 0.5).abs() < 1e-6, "got {got}"); + } + + #[test] + fn density_non_string_content_is_neutral() { + let mut m = Map::new(); + m.insert("role".to_string(), Value::String("assistant".to_string())); + m.insert( + "tool_calls".to_string(), + Value::Array(vec![Value::Object(Map::new())]), + ); + let v = Value::Object(m); + assert_eq!(MessageScorer::compute_density_score(&v), 0.5); + } + + #[test] + fn forward_refs_links_tool_response_to_assistant() { + let mut tc = Map::new(); + tc.insert("id".to_string(), Value::String("call_1".to_string())); + let mut assistant = Map::new(); + assistant.insert("role".to_string(), Value::String("assistant".to_string())); + assistant.insert( + "tool_calls".to_string(), + Value::Array(vec![Value::Object(tc)]), + ); + + let mut tool_resp = Map::new(); + tool_resp.insert("role".to_string(), Value::String("tool".to_string())); + tool_resp.insert( + "tool_call_id".to_string(), + Value::String("call_1".to_string()), + ); + tool_resp.insert("content".to_string(), Value::String("result".to_string())); + + let messages = vec![ + msg("user", "do thing"), + Value::Object(assistant), + Value::Object(tool_resp), + ]; + + let refs = MessageScorer::compute_forward_references(&messages); + assert_eq!(refs.get(&1), Some(&1)); + assert_eq!(refs.get(&0), None); + } + + #[test] + fn forward_refs_unmatched_tool_call_id_ignored() { + let mut tool_resp = Map::new(); + tool_resp.insert("role".to_string(), Value::String("tool".to_string())); + tool_resp.insert( + "tool_call_id".to_string(), + Value::String("nope".to_string()), + ); + tool_resp.insert("content".to_string(), Value::String("result".to_string())); + + let messages = vec![Value::Object(tool_resp)]; + let refs = MessageScorer::compute_forward_references(&messages); + assert!(refs.is_empty()); + } + + #[test] + fn reference_score_zero_refs_returns_zero() { + let s = MessageScorer::with_defaults(); + let refs = HashMap::new(); + assert_eq!(s.compute_reference_score(0, &refs), 0.0); + } + + #[test] + fn reference_score_one_ref_uses_log_formula() { + let s = MessageScorer::with_defaults(); + let mut refs = HashMap::new(); + refs.insert(0, 1); + // 0.3 + 0.2 * ln(2) = 0.3 + 0.2 * 0.693... ≈ 0.4386 + let expected = 0.3 + 0.2 * 2f32.ln(); + let got = s.compute_reference_score(0, &refs); + assert!((got - expected).abs() < 1e-6, "got {got}"); + } + + #[test] + fn reference_score_clamps_to_one() { + let s = MessageScorer::with_defaults(); + let mut refs = HashMap::new(); + refs.insert(0, 1_000_000); + assert!(s.compute_reference_score(0, &refs) <= 1.0); + } + + #[test] + fn semantic_score_no_provider_returns_neutral() { + let s = MessageScorer::with_defaults(); + let m = msg("user", "hello world"); + // recent_embedding is None too without a provider + assert_eq!(s.compute_semantic_score(&m, 0, None), 0.5); + } + + #[test] + fn toin_score_no_provider_returns_neutral() { + let s = MessageScorer::with_defaults(); + let m = msg("tool", "{}"); + assert_eq!(s.compute_toin_score(&m), 0.5); + } + + #[test] + fn error_score_no_provider_returns_zero() { + let s = MessageScorer::with_defaults(); + let m = msg("tool", "{}"); + assert_eq!(s.compute_error_score(&m), 0.0); + } + + #[test] + fn estimate_tokens_string_uses_char_count() { + // 16 chars / 4 = 4 + let m = msg("user", "abcdefghijklmnop"); + assert_eq!(estimate_tokens(&m), 4); + } + + #[test] + fn estimate_tokens_unicode_is_char_aware() { + // 4 emoji chars (each multi-byte) → len("...") in Python is 4. + let m = msg("user", "🎉🎉🎉🎉"); + // chars().count() = 4, // 4 = 1 + assert_eq!(estimate_tokens(&m), 1); + } + + #[test] + fn estimate_tokens_non_string_returns_default() { + let mut m = Map::new(); + m.insert("role".to_string(), Value::String("assistant".to_string())); + m.insert("tool_calls".to_string(), Value::Array(vec![])); + let v = Value::Object(m); + assert_eq!(estimate_tokens(&v), 100); + } + + #[test] + fn cosine_similarity_identical_vectors_is_one() { + let v = [1.0f32, 2.0, 3.0]; + assert!((cosine_similarity(&v, &v) - 1.0).abs() < 1e-6); + } + + #[test] + fn cosine_similarity_orthogonal_is_zero() { + let a = [1.0f32, 0.0]; + let b = [0.0f32, 1.0]; + assert!(cosine_similarity(&a, &b).abs() < 1e-6); + } + + #[test] + fn cosine_similarity_zero_vector_is_zero() { + let a = [0.0f32, 0.0]; + let b = [1.0f32, 1.0]; + assert_eq!(cosine_similarity(&a, &b), 0.0); + } + + #[test] + fn cosine_similarity_dim_mismatch_is_zero() { + let a = [1.0f32, 2.0]; + let b = [1.0f32, 2.0, 3.0]; + assert_eq!(cosine_similarity(&a, &b), 0.0); + } + + #[test] + fn drop_safe_mirrors_python_or_logic() { + // Python: drop_safe = not in_tool_unit OR not protected + // truth table: + // in_tool_unit=F, protected=F → T + // in_tool_unit=F, protected=T → T (not F = T) + // in_tool_unit=T, protected=F → T (not F = T) + // in_tool_unit=T, protected=T → F + let s = MessageScorer::with_defaults(); + let m = msg("user", "hi"); + let refs = HashMap::new(); + + let cases = [ + (false, false, true), + (false, true, true), + (true, false, true), + (true, true, false), + ]; + for (in_tu, prot, expected) in cases { + let score = s.score_message(&m, 0, 1, prot, in_tu, &refs, None); + assert_eq!( + score.drop_safe, expected, + "in_tool_unit={in_tu} protected={prot}" + ); + } + } + + #[test] + fn score_messages_returns_one_score_per_message() { + let s = MessageScorer::with_defaults(); + let messages = vec![ + msg("user", "hello"), + msg("assistant", "hi there"), + msg("user", "thanks"), + ]; + let scores = s.score_messages(&messages, &empty_set(), &empty_set()); + assert_eq!(scores.len(), 3); + assert_eq!(scores[0].message_index, 0); + assert_eq!(scores[2].message_index, 2); + } + + #[test] + fn score_messages_protected_indices_are_marked() { + let s = MessageScorer::with_defaults(); + let messages = vec![msg("user", "hi"), msg("user", "bye")]; + let mut protected = HashSet::new(); + protected.insert(0); + let scores = s.score_messages(&messages, &protected, &empty_set()); + assert!(scores[0].is_protected); + assert!(!scores[1].is_protected); + } + + #[test] + fn weights_are_normalized_on_construction() { + let unbalanced = ScoringWeights { + recency: 2.0, + semantic_similarity: 2.0, + toin_importance: 2.0, + error_indicator: 2.0, + forward_reference: 2.0, + token_density: 2.0, + }; + let s = MessageScorer::new(Some(unbalanced), None, None, 0.1); + // After normalization each should be 1/6. + assert!((s.weights.recency - 1.0 / 6.0).abs() < 1e-6); + } +} diff --git a/crates/headroom-core/src/scoring/traits.rs b/crates/headroom-core/src/scoring/traits.rs new file mode 100644 index 000000000..12327475f --- /dev/null +++ b/crates/headroom-core/src/scoring/traits.rs @@ -0,0 +1,175 @@ +//! External-dependency trait surfaces for `MessageScorer`. +//! +//! The scorer's six factors split into two groups: +//! +//! - **Pure / deterministic** (recency, forward references, density): +//! computed in-crate from the message list alone. No traits needed. +//! - **External-dependency** (semantic similarity, TOIN importance, +//! error indicator): require either an embedding model or learned +//! TOIN telemetry. These get trait surfaces here so the scorer +//! doesn't have to know how those subsystems are implemented. +//! +//! In PR-A (this PR), no concrete impls exist. PR-A1 wires +//! `EmbeddingProvider` to fastembed (reusing the `bge-small-en-v1.5` +//! model already loaded for SmartCrusher relevance). PR-A2 wires +//! `ToinProvider` to a PyO3 adapter so Rust can read Python's TOIN +//! state. When both land, the scorer code is unchanged — only the +//! provider construction at startup differs. +//! +//! # Why pass `&serde_json::Value` and not `&str` for tool content +//! +//! Python's `MessageScorer._compute_toin_score` parses the message +//! content as JSON, then computes a `ToolSignature` from the parsed +//! items, then looks up the learned pattern. The trait pushes the +//! parsing into the implementor — partly because the parsing is +//! cheap and message-local, partly because `ToolSignature` itself is +//! TOIN-internal (not yet ported, and not needed outside TOIN). + +use serde::{Deserialize, Serialize}; +use std::collections::BTreeMap; + +/// Embedding provider for semantic-similarity scoring. +/// +/// Implementations should return a fixed-dimension `Vec` for +/// each input text. The dimension must be consistent across calls +/// from a given instance — `MessageScorer` does cosine similarity +/// between vectors and assumes equal dimension. +/// +/// `embed` may fail in implementations that wrap external services; +/// the scorer treats failures as "no signal" (returns the neutral +/// `0.5` semantic score), matching Python's try/except behavior. +pub trait EmbeddingProvider: Send + Sync { + /// Embed a string into a fixed-dimension vector. Empty strings + /// or whitespace-only strings should still return a valid vector + /// — the scorer pre-filters empty content before calling this. + fn embed(&self, text: &str) -> Result, EmbeddingError>; +} + +/// Error type for embedding providers. Wraps the underlying impl's +/// error as a string for trait-object compatibility — losing typed +/// context is fine here since the scorer just logs + falls back. +#[derive(Debug, Clone)] +pub struct EmbeddingError(pub String); + +impl std::fmt::Display for EmbeddingError { + fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + write!(f, "embedding error: {}", self.0) + } +} + +impl std::error::Error for EmbeddingError {} + +/// TOIN (Tool Output Intelligence Network) pattern lookup. +/// +/// TOIN learns retrieval patterns per tool-output structure. The +/// scorer queries it for two things: +/// +/// 1. **Importance** (`toin_score`): high `retrieval_rate` means +/// users repeatedly retrieve this tool's data → keep it. +/// 2. **Error detection** (`error_score`): TOIN classifies fields +/// by inferred type. Fields tagged `error_indicator` boost the +/// error score, in lieu of hardcoded keyword regex. +/// +/// The trait takes the parsed JSON content directly (rather than a +/// pre-computed structure hash) because `ToolSignature` derivation +/// is TOIN-internal — pushing it across the trait boundary would +/// leak implementation details. +pub trait ToinProvider: Send + Sync { + /// Look up the learned pattern for a tool-message content + /// payload. Returns `None` if the content can't be classified + /// (not list/dict, empty, etc.) or no pattern has been learned + /// for this structure yet. + fn pattern_for_tool_content(&self, content: &serde_json::Value) -> Option; +} + +/// Snapshot of a TOIN-learned pattern for a single tool-output +/// structure. Mirrors the subset of `headroom.telemetry.toin.ToolPattern` +/// that `MessageScorer` actually reads. +/// +/// We don't mirror the full Python class — TOIN updates patterns +/// in-place during learning, but the scorer only reads. Returning +/// a snapshot decouples scorer reads from learning writes (no +/// shared-mutable-state across the FFI boundary in PR-A2). +#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)] +pub struct ToinPattern { + /// Overall confidence in this pattern, `[0.0, 1.0]`. Patterns + /// with `confidence < 0.3` are treated as not-yet-learned and + /// the scorer falls back to neutral. + pub confidence: f32, + + /// Fraction of tool invocations whose data was later retrieved + /// from cache or referenced by a follow-up message. `[0.0, 1.0]`. + /// High values mean "users keep needing this data" → important. + pub retrieval_rate: f32, + + /// Field hashes (NOT names — TOIN privacy-hashes them) that + /// are commonly retrieved from this structure. The scorer just + /// uses the *count* as a small importance boost; it doesn't + /// dereference individual hashes. + pub commonly_retrieved_fields: Vec, + + /// Per-field semantics, keyed by privacy-hashed field name. + /// `BTreeMap` rather than `HashMap` so serialization order is + /// deterministic (matters for parity-fixture stability). + pub field_semantics: BTreeMap, +} + +/// Inferred semantic type + confidence for a single field, as +/// learned by TOIN. Mirrors a subset of +/// `headroom.telemetry.models.FieldSemantics`. +#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)] +pub struct ToinFieldSemantic { + /// Inferred semantic category. The scorer specifically checks + /// for `"error_indicator"` to compute the error score; other + /// values (`"identifier"`, `"status"`, `"timestamp"`, etc.) are + /// not used by scoring but kept for completeness so we can pass + /// the same struct through to other consumers. + pub inferred_type: String, + + /// Confidence in the inferred type, `[0.0, 1.0]`. The scorer + /// applies a `>= 0.7` threshold for a "high-confidence error" + /// boost; below that it still counts the field but doesn't + /// boost. + pub confidence: f32, +} + +#[cfg(test)] +mod tests { + use super::*; + + #[test] + fn pattern_round_trips_through_serde() { + let mut field_semantics = BTreeMap::new(); + field_semantics.insert( + "abc123".to_string(), + ToinFieldSemantic { + inferred_type: "error_indicator".to_string(), + confidence: 0.85, + }, + ); + field_semantics.insert( + "def456".to_string(), + ToinFieldSemantic { + inferred_type: "identifier".to_string(), + confidence: 0.9, + }, + ); + + let p = ToinPattern { + confidence: 0.75, + retrieval_rate: 0.6, + commonly_retrieved_fields: vec!["abc123".to_string()], + field_semantics, + }; + + let json = serde_json::to_string(&p).unwrap(); + let back: ToinPattern = serde_json::from_str(&json).unwrap(); + assert_eq!(p, back); + } + + #[test] + fn embedding_error_is_displayable() { + let e = EmbeddingError("model not loaded".to_string()); + assert_eq!(e.to_string(), "embedding error: model not loaded"); + } +} diff --git a/crates/headroom-core/src/scoring/weights.rs b/crates/headroom-core/src/scoring/weights.rs new file mode 100644 index 000000000..a2b4a0524 --- /dev/null +++ b/crates/headroom-core/src/scoring/weights.rs @@ -0,0 +1,170 @@ +//! `ScoringWeights` — six-factor weighted importance scoring. +//! +//! Mirrors `headroom.config.ScoringWeights` byte-for-byte. The six +//! factors and their default contributions sum to ~1.0; `normalized()` +//! enforces that explicitly when a caller passes weights that don't +//! sum to 1.0 (e.g. learned weights from TOIN telemetry). +//! +//! Default values match Python's defaults exactly so the parity +//! fixtures byte-equal across implementations. + +use serde::{Deserialize, Serialize}; + +/// Weights for the six importance-scoring factors. +/// +/// All weights should sum to ~1.0 for normalized scoring (call +/// [`Self::normalized`] to enforce). Non-normalized weights are +/// permitted — `MessageScorer` does NOT auto-normalize on input — +/// since callers may use raw weights for relative comparison. +/// +/// Defaults match Python's `ScoringWeights()`: +/// +/// | Factor | Weight | +/// |--------|--------| +/// | recency | 0.20 | +/// | semantic_similarity | 0.20 | +/// | toin_importance | 0.25 | +/// | error_indicator | 0.15 | +/// | forward_reference | 0.15 | +/// | token_density | 0.05 | +#[derive(Debug, Clone, Copy, PartialEq, Serialize, Deserialize)] +pub struct ScoringWeights { + /// Exponential decay from conversation end. Default 0.20. + pub recency: f32, + /// Embedding cosine similarity to recent context. Default 0.20. + pub semantic_similarity: f32, + /// TOIN-learned field importance. Default 0.25. + pub toin_importance: f32, + /// TOIN-learned error-field detection. Default 0.15. + pub error_indicator: f32, + /// Number of later messages referencing this one (tool_call_id). Default 0.15. + pub forward_reference: f32, + /// Information density (unique-tokens / total-tokens). Default 0.05. + pub token_density: f32, +} + +impl Default for ScoringWeights { + fn default() -> Self { + Self { + recency: 0.20, + semantic_similarity: 0.20, + toin_importance: 0.25, + error_indicator: 0.15, + forward_reference: 0.15, + token_density: 0.05, + } + } +} + +impl ScoringWeights { + /// Return a copy with all weights divided by their sum, so they sum + /// to 1.0 exactly. If the input sums to 0 (degenerate config), + /// returns the default weights. + pub fn normalized(&self) -> Self { + let total = self.recency + + self.semantic_similarity + + self.toin_importance + + self.error_indicator + + self.forward_reference + + self.token_density; + if total == 0.0 { + return Self::default(); + } + Self { + recency: self.recency / total, + semantic_similarity: self.semantic_similarity / total, + toin_importance: self.toin_importance / total, + error_indicator: self.error_indicator / total, + forward_reference: self.forward_reference / total, + token_density: self.token_density / total, + } + } +} + +#[cfg(test)] +mod tests { + use super::*; + + #[test] + fn defaults_match_python() { + let w = ScoringWeights::default(); + assert_eq!(w.recency, 0.20); + assert_eq!(w.semantic_similarity, 0.20); + assert_eq!(w.toin_importance, 0.25); + assert_eq!(w.error_indicator, 0.15); + assert_eq!(w.forward_reference, 0.15); + assert_eq!(w.token_density, 0.05); + } + + #[test] + fn defaults_sum_to_one() { + let w = ScoringWeights::default(); + let total = w.recency + + w.semantic_similarity + + w.toin_importance + + w.error_indicator + + w.forward_reference + + w.token_density; + assert!((total - 1.0).abs() < 1e-6); + } + + #[test] + fn normalized_already_normal_is_idempotent() { + let w = ScoringWeights::default(); + let n = w.normalized(); + // Within float epsilon — divisions can introduce tiny drift. + assert!((n.recency - w.recency).abs() < 1e-6); + assert!((n.toin_importance - w.toin_importance).abs() < 1e-6); + } + + #[test] + fn normalized_unbalanced_weights_sum_to_one() { + let w = ScoringWeights { + recency: 1.0, + semantic_similarity: 1.0, + toin_importance: 1.0, + error_indicator: 1.0, + forward_reference: 1.0, + token_density: 1.0, + }; + let n = w.normalized(); + let total = n.recency + + n.semantic_similarity + + n.toin_importance + + n.error_indicator + + n.forward_reference + + n.token_density; + assert!((total - 1.0).abs() < 1e-6); + // Each component should now be ~1/6. + assert!((n.recency - 1.0 / 6.0).abs() < 1e-6); + } + + #[test] + fn normalized_zero_weights_falls_back_to_default() { + let w = ScoringWeights { + recency: 0.0, + semantic_similarity: 0.0, + toin_importance: 0.0, + error_indicator: 0.0, + forward_reference: 0.0, + token_density: 0.0, + }; + let n = w.normalized(); + assert_eq!(n, ScoringWeights::default()); + } + + #[test] + fn round_trips_through_serde() { + let w = ScoringWeights { + recency: 0.3, + semantic_similarity: 0.1, + toin_importance: 0.2, + error_indicator: 0.2, + forward_reference: 0.1, + token_density: 0.1, + }; + let json = serde_json::to_string(&w).unwrap(); + let back: ScoringWeights = serde_json::from_str(&json).unwrap(); + assert_eq!(w, back); + } +} diff --git a/crates/headroom-core/src/transforms/smart_crusher/crusher.rs b/crates/headroom-core/src/transforms/smart_crusher/crusher.rs index 39648132f..10efd5275 100644 --- a/crates/headroom-core/src/transforms/smart_crusher/crusher.rs +++ b/crates/headroom-core/src/transforms/smart_crusher/crusher.rs @@ -933,15 +933,6 @@ fn hash_canonical(canonical: &str) -> String { .collect() } -/// Convenience: canonical-serialize `items` and hash the result. Kept -/// for sites (e.g. tests) that don't also need the canonical bytes for -/// storage. Production lossy path inlines `canonical_array_json` + -/// `hash_canonical` so the bytes are reused for the store payload. -#[cfg(test)] -fn hash_array_for_ccr(items: &[Value]) -> String { - hash_canonical(&canonical_array_json(items)) -} - // ─── PR5 walker-integration helpers (string handling) ────────────────────── // // Parse-as-JSON-container, marker formatting, and humanize-bytes used to @@ -1277,8 +1268,10 @@ mod tests { // Use low-uniqueness items so the analyzer is willing to // crush (unique id+name per row would trip the // "unique_entities_no_signal" skip gate instead). - let mut cfg = SmartCrusherConfig::default(); - cfg.lossless_min_savings_ratio = 0.99; + let cfg = SmartCrusherConfig { + lossless_min_savings_ratio: 0.99, + ..Default::default() + }; let c = SmartCrusher::new(cfg); let items: Vec = (0..50).map(|_| json!({"status": "ok"})).collect(); let result = c.crush_array(&items, "", 1.0); @@ -1305,8 +1298,10 @@ mod tests { #[test] fn ccr_hash_is_deterministic() { // Same input → same hash, so the runtime cache key is stable. - let mut cfg = SmartCrusherConfig::default(); - cfg.lossless_min_savings_ratio = 0.99; // force lossy path + let cfg = SmartCrusherConfig { + lossless_min_savings_ratio: 0.99, // force lossy path + ..Default::default() + }; let c = SmartCrusher::new(cfg); let items: Vec = (0..30).map(|i| json!({"id": i, "tag": "ok"})).collect(); let r1 = c.crush_array(&items, "", 1.0); @@ -1317,8 +1312,10 @@ mod tests { #[test] fn ccr_hash_changes_with_input() { - let mut cfg = SmartCrusherConfig::default(); - cfg.lossless_min_savings_ratio = 0.99; + let cfg = SmartCrusherConfig { + lossless_min_savings_ratio: 0.99, + ..Default::default() + }; let c = SmartCrusher::new(cfg); let a: Vec = (0..30).map(|i| json!({"id": i})).collect(); let b: Vec = (100..130).map(|i| json!({"id": i})).collect(); diff --git a/crates/headroom-parity/src/lib.rs b/crates/headroom-parity/src/lib.rs index bf0b33e97..70e617d6b 100644 --- a/crates/headroom-parity/src/lib.rs +++ b/crates/headroom-parity/src/lib.rs @@ -462,6 +462,130 @@ impl TransformComparator for ContentDetectorComparator { } } +/// Real comparator for the `message_scorer` transform. Drives the +/// Rust port over fixture inputs and emits the same per-message +/// score struct shape the Python recorder dumps. +/// +/// Float parity strategy: deterministic factors use `f32::exp` / +/// `f32::ln` on the Rust side and `math.exp(f64)` / `math.log(f64)` +/// on the Python side. Both are mathematically identical but their +/// last-bit rounding can differ. The Python recorder rounds every +/// float to 6 decimals before writing the fixture; this comparator +/// mirrors the same rounding so JSON byte-equality holds. +pub struct MessageScorerComparator; + +impl TransformComparator for MessageScorerComparator { + fn name(&self) -> &str { + "message_scorer" + } + + fn run( + &self, + input: &serde_json::Value, + config: &serde_json::Value, + ) -> Result { + use headroom_core::scoring::{MessageScorer, ScoringWeights}; + use std::collections::HashSet; + + let messages: Vec = input + .get("messages") + .and_then(|v| v.as_array()) + .cloned() + .context("message_scorer fixture input.messages must be an array")?; + + let parse_index_set = |key: &str| -> HashSet { + input + .get(key) + .and_then(|v| v.as_array()) + .map(|arr| { + arr.iter() + .filter_map(|v| v.as_u64().map(|n| n as usize)) + .collect() + }) + .unwrap_or_default() + }; + let protected = parse_index_set("protected_indices"); + let tool_unit = parse_index_set("tool_unit_indices"); + + let decay_rate = input + .get("decay_rate") + .and_then(|v| v.as_f64()) + .map(|v| v as f32) + .unwrap_or(0.1); + + // config.weights may be null (use defaults) or a full struct. + let weights = config + .get("weights") + .and_then(|v| if v.is_null() { None } else { Some(v) }) + .map(|w| ScoringWeights { + recency: w.get("recency").and_then(|x| x.as_f64()).unwrap_or(0.20) as f32, + semantic_similarity: w + .get("semantic_similarity") + .and_then(|x| x.as_f64()) + .unwrap_or(0.20) as f32, + toin_importance: w + .get("toin_importance") + .and_then(|x| x.as_f64()) + .unwrap_or(0.25) as f32, + error_indicator: w + .get("error_indicator") + .and_then(|x| x.as_f64()) + .unwrap_or(0.15) as f32, + forward_reference: w + .get("forward_reference") + .and_then(|x| x.as_f64()) + .unwrap_or(0.15) as f32, + token_density: w + .get("token_density") + .and_then(|x| x.as_f64()) + .unwrap_or(0.05) as f32, + }); + + let scorer = MessageScorer::new(weights, None, None, decay_rate); + let scores = scorer.score_messages(&messages, &protected, &tool_unit); + + let scored_array: Vec = scores + .into_iter() + .map(|s| serde_json::to_value(&s).expect("MessageScore is serializable")) + .collect(); + + // 5 decimal places matches the Python recorder. See + // record_message_scorer.py:_FLOAT_ROUND_PLACES for why. + Ok(round_floats(&serde_json::Value::Array(scored_array), 5)) + } +} + +/// Recursively round every f64 value in a JSON tree to `places` +/// decimal places. Used to absorb f32-vs-f64 last-bit drift between +/// the Python recorder and Rust comparator. +fn round_floats(value: &serde_json::Value, places: u32) -> serde_json::Value { + match value { + serde_json::Value::Number(n) => { + if let Some(f) = n.as_f64() { + if f.is_finite() && n.is_f64() { + let factor = 10f64.powi(places as i32); + let rounded = (f * factor).round() / factor; + return serde_json::Number::from_f64(rounded) + .map(serde_json::Value::Number) + .unwrap_or_else(|| value.clone()); + } + } + value.clone() + } + serde_json::Value::Array(arr) => { + serde_json::Value::Array(arr.iter().map(|v| round_floats(v, places)).collect()) + } + serde_json::Value::Object(map) => { + let mut out = serde_json::Map::new(); + for (k, v) in map { + out.insert(k.clone(), round_floats(v, places)); + } + serde_json::Value::Object(out) + } + _ => value.clone(), + } +} + /// Every built-in comparator, in a stable order. pub fn builtin_comparators() -> Vec> { vec![ @@ -472,6 +596,7 @@ pub fn builtin_comparators() -> Vec> { Box::new(CcrComparator), Box::new(SmartCrusherComparator), Box::new(ContentDetectorComparator), + Box::new(MessageScorerComparator), ] } diff --git a/tests/parity/fixtures/message_scorer/custom_weights_recency_heavy_5cc105616d48.json b/tests/parity/fixtures/message_scorer/custom_weights_recency_heavy_5cc105616d48.json new file mode 100644 index 000000000..1a41b6c20 --- /dev/null +++ b/tests/parity/fixtures/message_scorer/custom_weights_recency_heavy_5cc105616d48.json @@ -0,0 +1,125 @@ +{ + "config": { + "weights": { + "error_indicator": 0.05, + "forward_reference": 0.1, + "recency": 0.6, + "semantic_similarity": 0.1, + "toin_importance": 0.1, + "token_density": 0.05 + } + }, + "input": { + "decay_rate": 0.1, + "messages": [ + { + "content": "first message in a longer chat", + "role": "user" + }, + { + "content": "an assistant reply with substance", + "role": "assistant" + }, + { + "content": "another follow up question here", + "role": "user" + }, + { + "content": "and the closing reply", + "role": "assistant" + } + ], + "protected_indices": [], + "tool_unit_indices": [] + }, + "input_sha256": "5cc105616d481b3c37f1cbf359b16967669b1703acc83c7df7e70488502271a6", + "label": "custom_weights_recency_heavy", + "output": [ + { + "density_score": 1.0, + "drop_safe": true, + "error_score": 0.0, + "is_protected": false, + "message_index": 0, + "recency_score": 0.74082, + "reference_score": 0.0, + "score_breakdown": { + "density": 1.0, + "error": 0.0, + "recency": 0.74082, + "reference": 0.0, + "semantic": 0.5, + "toin": 0.5 + }, + "semantic_score": 0.5, + "toin_score": 0.5, + "tokens": 7, + "total_score": 0.59449 + }, + { + "density_score": 1.0, + "drop_safe": true, + "error_score": 0.0, + "is_protected": false, + "message_index": 1, + "recency_score": 0.81873, + "reference_score": 0.0, + "score_breakdown": { + "density": 1.0, + "error": 0.0, + "recency": 0.81873, + "reference": 0.0, + "semantic": 0.5, + "toin": 0.5 + }, + "semantic_score": 0.5, + "toin_score": 0.5, + "tokens": 8, + "total_score": 0.64124 + }, + { + "density_score": 1.0, + "drop_safe": true, + "error_score": 0.0, + "is_protected": false, + "message_index": 2, + "recency_score": 0.90484, + "reference_score": 0.0, + "score_breakdown": { + "density": 1.0, + "error": 0.0, + "recency": 0.90484, + "reference": 0.0, + "semantic": 0.5, + "toin": 0.5 + }, + "semantic_score": 0.5, + "toin_score": 0.5, + "tokens": 7, + "total_score": 0.6929 + }, + { + "density_score": 1.0, + "drop_safe": true, + "error_score": 0.0, + "is_protected": false, + "message_index": 3, + "recency_score": 1.0, + "reference_score": 0.0, + "score_breakdown": { + "density": 1.0, + "error": 0.0, + "recency": 1.0, + "reference": 0.0, + "semantic": 0.5, + "toin": 0.5 + }, + "semantic_score": 0.5, + "toin_score": 0.5, + "tokens": 5, + "total_score": 0.75 + } + ], + "recorded_at": "2026-05-01T04:29:36.075652+00:00", + "transform": "message_scorer" +} diff --git a/tests/parity/fixtures/message_scorer/drop_safe_truth_table_2e03ebdc1b01.json b/tests/parity/fixtures/message_scorer/drop_safe_truth_table_2e03ebdc1b01.json new file mode 100644 index 000000000..ee1970287 --- /dev/null +++ b/tests/parity/fixtures/message_scorer/drop_safe_truth_table_2e03ebdc1b01.json @@ -0,0 +1,133 @@ +{ + "config": { + "weights": null + }, + "input": { + "decay_rate": 0.1, + "messages": [ + { + "content": "you are helpful", + "role": "system" + }, + { + "content": "hello", + "role": "user" + }, + { + "content": "", + "role": "assistant", + "tool_calls": [ + { + "function": { + "name": "f" + }, + "id": "x" + } + ] + }, + { + "content": "result", + "role": "tool", + "tool_call_id": "x" + } + ], + "protected_indices": [ + 0, + 2 + ], + "tool_unit_indices": [ + 2, + 3 + ] + }, + "input_sha256": "2e03ebdc1b0185fdd33c537f199890e5faabda24a4ee3025959419c21eda5569", + "label": "drop_safe_truth_table", + "output": [ + { + "density_score": 1.0, + "drop_safe": true, + "error_score": 0.0, + "is_protected": true, + "message_index": 0, + "recency_score": 0.74082, + "reference_score": 0.0, + "score_breakdown": { + "density": 1.0, + "error": 0.0, + "recency": 0.74082, + "reference": 0.0, + "semantic": 0.5, + "toin": 0.5 + }, + "semantic_score": 0.5, + "toin_score": 0.5, + "tokens": 3, + "total_score": 0.42316 + }, + { + "density_score": 0.5, + "drop_safe": true, + "error_score": 0.0, + "is_protected": false, + "message_index": 1, + "recency_score": 0.81873, + "reference_score": 0.0, + "score_breakdown": { + "density": 0.5, + "error": 0.0, + "recency": 0.81873, + "reference": 0.0, + "semantic": 0.5, + "toin": 0.5 + }, + "semantic_score": 0.5, + "toin_score": 0.5, + "tokens": 1, + "total_score": 0.41375 + }, + { + "density_score": 0.5, + "drop_safe": false, + "error_score": 0.0, + "is_protected": true, + "message_index": 2, + "recency_score": 0.90484, + "reference_score": 0.43863, + "score_breakdown": { + "density": 0.5, + "error": 0.0, + "recency": 0.90484, + "reference": 0.43863, + "semantic": 0.5, + "toin": 0.5 + }, + "semantic_score": 0.5, + "toin_score": 0.5, + "tokens": 0, + "total_score": 0.49676 + }, + { + "density_score": 0.5, + "drop_safe": true, + "error_score": 0.0, + "is_protected": false, + "message_index": 3, + "recency_score": 1.0, + "reference_score": 0.0, + "score_breakdown": { + "density": 0.5, + "error": 0.0, + "recency": 1.0, + "reference": 0.0, + "semantic": 0.5, + "toin": 0.5 + }, + "semantic_score": 0.5, + "toin_score": 0.5, + "tokens": 1, + "total_score": 0.45 + } + ], + "recorded_at": "2026-05-01T04:29:36.076447+00:00", + "transform": "message_scorer" +} diff --git a/tests/parity/fixtures/message_scorer/empty_c94d65311b40.json b/tests/parity/fixtures/message_scorer/empty_c94d65311b40.json new file mode 100644 index 000000000..cadcb7156 --- /dev/null +++ b/tests/parity/fixtures/message_scorer/empty_c94d65311b40.json @@ -0,0 +1,16 @@ +{ + "config": { + "weights": null + }, + "input": { + "decay_rate": 0.1, + "messages": [], + "protected_indices": [], + "tool_unit_indices": [] + }, + "input_sha256": "c94d65311b40eb4065d7160bded90e79fd48aebca28fc334c257ea658c66413b", + "label": "empty", + "output": [], + "recorded_at": "2026-05-01T04:29:36.073767+00:00", + "transform": "message_scorer" +} diff --git a/tests/parity/fixtures/message_scorer/high_density_dd454e2933e3.json b/tests/parity/fixtures/message_scorer/high_density_dd454e2933e3.json new file mode 100644 index 000000000..05cc49e4b --- /dev/null +++ b/tests/parity/fixtures/message_scorer/high_density_dd454e2933e3.json @@ -0,0 +1,43 @@ +{ + "config": { + "weights": null + }, + "input": { + "decay_rate": 0.1, + "messages": [ + { + "content": "alpha bravo charlie delta echo foxtrot golf hotel", + "role": "user" + } + ], + "protected_indices": [], + "tool_unit_indices": [] + }, + "input_sha256": "dd454e2933e3c2d10ab2fdfbb4675f797711cbd2e67d44ed767a9c590121c88c", + "label": "high_density", + "output": [ + { + "density_score": 1.0, + "drop_safe": true, + "error_score": 0.0, + "is_protected": false, + "message_index": 0, + "recency_score": 1.0, + "reference_score": 0.0, + "score_breakdown": { + "density": 1.0, + "error": 0.0, + "recency": 1.0, + "reference": 0.0, + "semantic": 0.5, + "toin": 0.5 + }, + "semantic_score": 0.5, + "toin_score": 0.5, + "tokens": 12, + "total_score": 0.475 + } + ], + "recorded_at": "2026-05-01T04:29:36.075245+00:00", + "transform": "message_scorer" +} diff --git a/tests/parity/fixtures/message_scorer/linear_5_messages_71f1ec6a317c.json b/tests/parity/fixtures/message_scorer/linear_5_messages_71f1ec6a317c.json new file mode 100644 index 000000000..b0c071079 --- /dev/null +++ b/tests/parity/fixtures/message_scorer/linear_5_messages_71f1ec6a317c.json @@ -0,0 +1,145 @@ +{ + "config": { + "weights": null + }, + "input": { + "decay_rate": 0.1, + "messages": [ + { + "content": "what is the capital of france", + "role": "user" + }, + { + "content": "the capital of france is paris", + "role": "assistant" + }, + { + "content": "and germany", + "role": "user" + }, + { + "content": "the capital of germany is berlin", + "role": "assistant" + }, + { + "content": "thanks", + "role": "user" + } + ], + "protected_indices": [ + 0 + ], + "tool_unit_indices": [] + }, + "input_sha256": "71f1ec6a317cdf54004ec850d9d7312c7eed622337d3961fab83cb56c1f8b702", + "label": "linear_5_messages", + "output": [ + { + "density_score": 1.0, + "drop_safe": true, + "error_score": 0.0, + "is_protected": true, + "message_index": 0, + "recency_score": 0.67032, + "reference_score": 0.0, + "score_breakdown": { + "density": 1.0, + "error": 0.0, + "recency": 0.67032, + "reference": 0.0, + "semantic": 0.5, + "toin": 0.5 + }, + "semantic_score": 0.5, + "toin_score": 0.5, + "tokens": 7, + "total_score": 0.40906 + }, + { + "density_score": 1.0, + "drop_safe": true, + "error_score": 0.0, + "is_protected": false, + "message_index": 1, + "recency_score": 0.74082, + "reference_score": 0.0, + "score_breakdown": { + "density": 1.0, + "error": 0.0, + "recency": 0.74082, + "reference": 0.0, + "semantic": 0.5, + "toin": 0.5 + }, + "semantic_score": 0.5, + "toin_score": 0.5, + "tokens": 7, + "total_score": 0.42316 + }, + { + "density_score": 0.5, + "drop_safe": true, + "error_score": 0.0, + "is_protected": false, + "message_index": 2, + "recency_score": 0.81873, + "reference_score": 0.0, + "score_breakdown": { + "density": 0.5, + "error": 0.0, + "recency": 0.81873, + "reference": 0.0, + "semantic": 0.5, + "toin": 0.5 + }, + "semantic_score": 0.5, + "toin_score": 0.5, + "tokens": 2, + "total_score": 0.41375 + }, + { + "density_score": 1.0, + "drop_safe": true, + "error_score": 0.0, + "is_protected": false, + "message_index": 3, + "recency_score": 0.90484, + "reference_score": 0.0, + "score_breakdown": { + "density": 1.0, + "error": 0.0, + "recency": 0.90484, + "reference": 0.0, + "semantic": 0.5, + "toin": 0.5 + }, + "semantic_score": 0.5, + "toin_score": 0.5, + "tokens": 8, + "total_score": 0.45597 + }, + { + "density_score": 0.5, + "drop_safe": true, + "error_score": 0.0, + "is_protected": false, + "message_index": 4, + "recency_score": 1.0, + "reference_score": 0.0, + "score_breakdown": { + "density": 0.5, + "error": 0.0, + "recency": 1.0, + "reference": 0.0, + "semantic": 0.5, + "toin": 0.5 + }, + "semantic_score": 0.5, + "toin_score": 0.5, + "tokens": 1, + "total_score": 0.45 + } + ], + "recorded_at": "2026-05-01T04:29:36.074114+00:00", + "transform": "message_scorer" +} diff --git a/tests/parity/fixtures/message_scorer/low_density_28ad2371fdc0.json b/tests/parity/fixtures/message_scorer/low_density_28ad2371fdc0.json new file mode 100644 index 000000000..eee09fa8c --- /dev/null +++ b/tests/parity/fixtures/message_scorer/low_density_28ad2371fdc0.json @@ -0,0 +1,43 @@ +{ + "config": { + "weights": null + }, + "input": { + "decay_rate": 0.1, + "messages": [ + { + "content": "ok ok ok ok ok ok ok ok ok ok", + "role": "user" + } + ], + "protected_indices": [], + "tool_unit_indices": [] + }, + "input_sha256": "28ad2371fdc08517452da84c371a2063fff941f2613c87cafef8c70efa8bf4ba", + "label": "low_density", + "output": [ + { + "density_score": 0.0, + "drop_safe": true, + "error_score": 0.0, + "is_protected": false, + "message_index": 0, + "recency_score": 1.0, + "reference_score": 0.0, + "score_breakdown": { + "density": 0.0, + "error": 0.0, + "recency": 1.0, + "reference": 0.0, + "semantic": 0.5, + "toin": 0.5 + }, + "semantic_score": 0.5, + "toin_score": 0.5, + "tokens": 7, + "total_score": 0.425 + } + ], + "recorded_at": "2026-05-01T04:29:36.075414+00:00", + "transform": "message_scorer" +} diff --git a/tests/parity/fixtures/message_scorer/multi_ref_to_assistant_67a047a411c4.json b/tests/parity/fixtures/message_scorer/multi_ref_to_assistant_67a047a411c4.json new file mode 100644 index 000000000..29a9d6c73 --- /dev/null +++ b/tests/parity/fixtures/message_scorer/multi_ref_to_assistant_67a047a411c4.json @@ -0,0 +1,171 @@ +{ + "config": { + "weights": null + }, + "input": { + "decay_rate": 0.1, + "messages": [ + { + "content": "do many things", + "role": "user" + }, + { + "content": "", + "role": "assistant", + "tool_calls": [ + { + "function": { + "name": "f" + }, + "id": "c1" + }, + { + "function": { + "name": "g" + }, + "id": "c2" + }, + { + "function": { + "name": "h" + }, + "id": "c3" + } + ] + }, + { + "content": "r1", + "role": "tool", + "tool_call_id": "c1" + }, + { + "content": "r2", + "role": "tool", + "tool_call_id": "c2" + }, + { + "content": "r3", + "role": "tool", + "tool_call_id": "c3" + } + ], + "protected_indices": [], + "tool_unit_indices": [ + 1, + 2, + 3, + 4 + ] + }, + "input_sha256": "67a047a411c474d0d3fa712b89aabc12311910bd466f3c750c003245558306de", + "label": "multi_ref_to_assistant", + "output": [ + { + "density_score": 1.0, + "drop_safe": true, + "error_score": 0.0, + "is_protected": false, + "message_index": 0, + "recency_score": 0.67032, + "reference_score": 0.0, + "score_breakdown": { + "density": 1.0, + "error": 0.0, + "recency": 0.67032, + "reference": 0.0, + "semantic": 0.5, + "toin": 0.5 + }, + "semantic_score": 0.5, + "toin_score": 0.5, + "tokens": 3, + "total_score": 0.40906 + }, + { + "density_score": 0.5, + "drop_safe": true, + "error_score": 0.0, + "is_protected": false, + "message_index": 1, + "recency_score": 0.74082, + "reference_score": 0.57726, + "score_breakdown": { + "density": 0.5, + "error": 0.0, + "recency": 0.74082, + "reference": 0.57726, + "semantic": 0.5, + "toin": 0.5 + }, + "semantic_score": 0.5, + "toin_score": 0.5, + "tokens": 0, + "total_score": 0.48475 + }, + { + "density_score": 0.5, + "drop_safe": true, + "error_score": 0.0, + "is_protected": false, + "message_index": 2, + "recency_score": 0.81873, + "reference_score": 0.0, + "score_breakdown": { + "density": 0.5, + "error": 0.0, + "recency": 0.81873, + "reference": 0.0, + "semantic": 0.5, + "toin": 0.5 + }, + "semantic_score": 0.5, + "toin_score": 0.5, + "tokens": 0, + "total_score": 0.41375 + }, + { + "density_score": 0.5, + "drop_safe": true, + "error_score": 0.0, + "is_protected": false, + "message_index": 3, + "recency_score": 0.90484, + "reference_score": 0.0, + "score_breakdown": { + "density": 0.5, + "error": 0.0, + "recency": 0.90484, + "reference": 0.0, + "semantic": 0.5, + "toin": 0.5 + }, + "semantic_score": 0.5, + "toin_score": 0.5, + "tokens": 0, + "total_score": 0.43097 + }, + { + "density_score": 0.5, + "drop_safe": true, + "error_score": 0.0, + "is_protected": false, + "message_index": 4, + "recency_score": 1.0, + "reference_score": 0.0, + "score_breakdown": { + "density": 0.5, + "error": 0.0, + "recency": 1.0, + "reference": 0.0, + "semantic": 0.5, + "toin": 0.5 + }, + "semantic_score": 0.5, + "toin_score": 0.5, + "tokens": 0, + "total_score": 0.45 + } + ], + "recorded_at": "2026-05-01T04:29:36.074987+00:00", + "transform": "message_scorer" +} diff --git a/tests/parity/fixtures/message_scorer/non_string_content_9ab4060e7b90.json b/tests/parity/fixtures/message_scorer/non_string_content_9ab4060e7b90.json new file mode 100644 index 000000000..7f8db7841 --- /dev/null +++ b/tests/parity/fixtures/message_scorer/non_string_content_9ab4060e7b90.json @@ -0,0 +1,76 @@ +{ + "config": { + "weights": null + }, + "input": { + "decay_rate": 0.1, + "messages": [ + { + "content": "hi", + "role": "user" + }, + { + "content": null, + "role": "assistant", + "tool_calls": [ + { + "function": { + "name": "f" + }, + "id": "tc" + } + ] + } + ], + "protected_indices": [], + "tool_unit_indices": [] + }, + "input_sha256": "9ab4060e7b906c73fae6926ae4766fc651b1ab73fdf023496f1b9bd6d84c2cae", + "label": "non_string_content", + "output": [ + { + "density_score": 0.5, + "drop_safe": true, + "error_score": 0.0, + "is_protected": false, + "message_index": 0, + "recency_score": 0.90484, + "reference_score": 0.0, + "score_breakdown": { + "density": 0.5, + "error": 0.0, + "recency": 0.90484, + "reference": 0.0, + "semantic": 0.5, + "toin": 0.5 + }, + "semantic_score": 0.5, + "toin_score": 0.5, + "tokens": 0, + "total_score": 0.43097 + }, + { + "density_score": 0.5, + "drop_safe": true, + "error_score": 0.0, + "is_protected": false, + "message_index": 1, + "recency_score": 1.0, + "reference_score": 0.0, + "score_breakdown": { + "density": 0.5, + "error": 0.0, + "recency": 1.0, + "reference": 0.0, + "semantic": 0.5, + "toin": 0.5 + }, + "semantic_score": 0.5, + "toin_score": 0.5, + "tokens": 100, + "total_score": 0.45 + } + ], + "recorded_at": "2026-05-01T04:29:36.077209+00:00", + "transform": "message_scorer" +} diff --git a/tests/parity/fixtures/message_scorer/orphan_tool_response_adb4e13b7541.json b/tests/parity/fixtures/message_scorer/orphan_tool_response_adb4e13b7541.json new file mode 100644 index 000000000..45ac97627 --- /dev/null +++ b/tests/parity/fixtures/message_scorer/orphan_tool_response_adb4e13b7541.json @@ -0,0 +1,69 @@ +{ + "config": { + "weights": null + }, + "input": { + "decay_rate": 0.1, + "messages": [ + { + "content": "go", + "role": "user" + }, + { + "content": "stranded", + "role": "tool", + "tool_call_id": "nope" + } + ], + "protected_indices": [], + "tool_unit_indices": [] + }, + "input_sha256": "adb4e13b75416cfbf8623002eaf6446ee67060bc64f88182b06f3b431c140c2f", + "label": "orphan_tool_response", + "output": [ + { + "density_score": 0.5, + "drop_safe": true, + "error_score": 0.0, + "is_protected": false, + "message_index": 0, + "recency_score": 0.90484, + "reference_score": 0.0, + "score_breakdown": { + "density": 0.5, + "error": 0.0, + "recency": 0.90484, + "reference": 0.0, + "semantic": 0.5, + "toin": 0.5 + }, + "semantic_score": 0.5, + "toin_score": 0.5, + "tokens": 0, + "total_score": 0.43097 + }, + { + "density_score": 0.5, + "drop_safe": true, + "error_score": 0.0, + "is_protected": false, + "message_index": 1, + "recency_score": 1.0, + "reference_score": 0.0, + "score_breakdown": { + "density": 0.5, + "error": 0.0, + "recency": 1.0, + "reference": 0.0, + "semantic": 0.5, + "toin": 0.5 + }, + "semantic_score": 0.5, + "toin_score": 0.5, + "tokens": 2, + "total_score": 0.45 + } + ], + "recorded_at": "2026-05-01T04:29:36.077001+00:00", + "transform": "message_scorer" +} diff --git a/tests/parity/fixtures/message_scorer/single_user_message_2b4fac2277be.json b/tests/parity/fixtures/message_scorer/single_user_message_2b4fac2277be.json new file mode 100644 index 000000000..5931893e5 --- /dev/null +++ b/tests/parity/fixtures/message_scorer/single_user_message_2b4fac2277be.json @@ -0,0 +1,43 @@ +{ + "config": { + "weights": null + }, + "input": { + "decay_rate": 0.1, + "messages": [ + { + "content": "hello world", + "role": "user" + } + ], + "protected_indices": [], + "tool_unit_indices": [] + }, + "input_sha256": "2b4fac2277be4cee6dcb022904f561d8ea4753a82714afe829f9203e1c3e441e", + "label": "single_user_message", + "output": [ + { + "density_score": 0.5, + "drop_safe": true, + "error_score": 0.0, + "is_protected": false, + "message_index": 0, + "recency_score": 1.0, + "reference_score": 0.0, + "score_breakdown": { + "density": 0.5, + "error": 0.0, + "recency": 1.0, + "reference": 0.0, + "semantic": 0.5, + "toin": 0.5 + }, + "semantic_score": 0.5, + "toin_score": 0.5, + "tokens": 2, + "total_score": 0.45 + } + ], + "recorded_at": "2026-05-01T04:29:36.073288+00:00", + "transform": "message_scorer" +} diff --git a/tests/parity/fixtures/message_scorer/slow_decay_de4fafb3a42b.json b/tests/parity/fixtures/message_scorer/slow_decay_de4fafb3a42b.json new file mode 100644 index 000000000..b68e41eee --- /dev/null +++ b/tests/parity/fixtures/message_scorer/slow_decay_de4fafb3a42b.json @@ -0,0 +1,218 @@ +{ + "config": { + "weights": null + }, + "input": { + "decay_rate": 0.02, + "messages": [ + { + "content": "message number 0", + "role": "user" + }, + { + "content": "message number 1", + "role": "user" + }, + { + "content": "message number 2", + "role": "user" + }, + { + "content": "message number 3", + "role": "user" + }, + { + "content": "message number 4", + "role": "user" + }, + { + "content": "message number 5", + "role": "user" + }, + { + "content": "message number 6", + "role": "user" + }, + { + "content": "message number 7", + "role": "user" + } + ], + "protected_indices": [], + "tool_unit_indices": [] + }, + "input_sha256": "de4fafb3a42bca1ff0ec75d0530b41b40d5d5bfb424975beae5ad9099264112c", + "label": "slow_decay", + "output": [ + { + "density_score": 1.0, + "drop_safe": true, + "error_score": 0.0, + "is_protected": false, + "message_index": 0, + "recency_score": 0.86936, + "reference_score": 0.0, + "score_breakdown": { + "density": 1.0, + "error": 0.0, + "recency": 0.86936, + "reference": 0.0, + "semantic": 0.5, + "toin": 0.5 + }, + "semantic_score": 0.5, + "toin_score": 0.5, + "tokens": 4, + "total_score": 0.44887 + }, + { + "density_score": 1.0, + "drop_safe": true, + "error_score": 0.0, + "is_protected": false, + "message_index": 1, + "recency_score": 0.88692, + "reference_score": 0.0, + "score_breakdown": { + "density": 1.0, + "error": 0.0, + "recency": 0.88692, + "reference": 0.0, + "semantic": 0.5, + "toin": 0.5 + }, + "semantic_score": 0.5, + "toin_score": 0.5, + "tokens": 4, + "total_score": 0.45238 + }, + { + "density_score": 1.0, + "drop_safe": true, + "error_score": 0.0, + "is_protected": false, + "message_index": 2, + "recency_score": 0.90484, + "reference_score": 0.0, + "score_breakdown": { + "density": 1.0, + "error": 0.0, + "recency": 0.90484, + "reference": 0.0, + "semantic": 0.5, + "toin": 0.5 + }, + "semantic_score": 0.5, + "toin_score": 0.5, + "tokens": 4, + "total_score": 0.45597 + }, + { + "density_score": 1.0, + "drop_safe": true, + "error_score": 0.0, + "is_protected": false, + "message_index": 3, + "recency_score": 0.92312, + "reference_score": 0.0, + "score_breakdown": { + "density": 1.0, + "error": 0.0, + "recency": 0.92312, + "reference": 0.0, + "semantic": 0.5, + "toin": 0.5 + }, + "semantic_score": 0.5, + "toin_score": 0.5, + "tokens": 4, + "total_score": 0.45962 + }, + { + "density_score": 1.0, + "drop_safe": true, + "error_score": 0.0, + "is_protected": false, + "message_index": 4, + "recency_score": 0.94176, + "reference_score": 0.0, + "score_breakdown": { + "density": 1.0, + "error": 0.0, + "recency": 0.94176, + "reference": 0.0, + "semantic": 0.5, + "toin": 0.5 + }, + "semantic_score": 0.5, + "toin_score": 0.5, + "tokens": 4, + "total_score": 0.46335 + }, + { + "density_score": 1.0, + "drop_safe": true, + "error_score": 0.0, + "is_protected": false, + "message_index": 5, + "recency_score": 0.96079, + "reference_score": 0.0, + "score_breakdown": { + "density": 1.0, + "error": 0.0, + "recency": 0.96079, + "reference": 0.0, + "semantic": 0.5, + "toin": 0.5 + }, + "semantic_score": 0.5, + "toin_score": 0.5, + "tokens": 4, + "total_score": 0.46716 + }, + { + "density_score": 1.0, + "drop_safe": true, + "error_score": 0.0, + "is_protected": false, + "message_index": 6, + "recency_score": 0.9802, + "reference_score": 0.0, + "score_breakdown": { + "density": 1.0, + "error": 0.0, + "recency": 0.9802, + "reference": 0.0, + "semantic": 0.5, + "toin": 0.5 + }, + "semantic_score": 0.5, + "toin_score": 0.5, + "tokens": 4, + "total_score": 0.47104 + }, + { + "density_score": 1.0, + "drop_safe": true, + "error_score": 0.0, + "is_protected": false, + "message_index": 7, + "recency_score": 1.0, + "reference_score": 0.0, + "score_breakdown": { + "density": 1.0, + "error": 0.0, + "recency": 1.0, + "reference": 0.0, + "semantic": 0.5, + "toin": 0.5 + }, + "semantic_score": 0.5, + "toin_score": 0.5, + "tokens": 4, + "total_score": 0.475 + } + ], + "recorded_at": "2026-05-01T04:29:36.076024+00:00", + "transform": "message_scorer" +} diff --git a/tests/parity/fixtures/message_scorer/tool_call_pair_3443697600a6.json b/tests/parity/fixtures/message_scorer/tool_call_pair_3443697600a6.json new file mode 100644 index 000000000..519d409b1 --- /dev/null +++ b/tests/parity/fixtures/message_scorer/tool_call_pair_3443697600a6.json @@ -0,0 +1,134 @@ +{ + "config": { + "weights": null + }, + "input": { + "decay_rate": 0.1, + "messages": [ + { + "content": "what's the weather", + "role": "user" + }, + { + "content": "", + "role": "assistant", + "tool_calls": [ + { + "function": { + "arguments": "{}", + "name": "get_weather" + }, + "id": "call_1", + "type": "function" + } + ] + }, + { + "content": "{\"temp\": 72, \"conditions\": \"sunny\"}", + "role": "tool", + "tool_call_id": "call_1" + }, + { + "content": "it is 72 and sunny", + "role": "assistant" + } + ], + "protected_indices": [ + 0 + ], + "tool_unit_indices": [ + 1, + 2 + ] + }, + "input_sha256": "3443697600a6878b18e41dd4f499829ee9f0ea882cf0510b030e15ce1febd07d", + "label": "tool_call_pair", + "output": [ + { + "density_score": 1.0, + "drop_safe": true, + "error_score": 0.0, + "is_protected": true, + "message_index": 0, + "recency_score": 0.74082, + "reference_score": 0.0, + "score_breakdown": { + "density": 1.0, + "error": 0.0, + "recency": 0.74082, + "reference": 0.0, + "semantic": 0.5, + "toin": 0.5 + }, + "semantic_score": 0.5, + "toin_score": 0.5, + "tokens": 4, + "total_score": 0.42316 + }, + { + "density_score": 0.5, + "drop_safe": true, + "error_score": 0.0, + "is_protected": false, + "message_index": 1, + "recency_score": 0.81873, + "reference_score": 0.43863, + "score_breakdown": { + "density": 0.5, + "error": 0.0, + "recency": 0.81873, + "reference": 0.43863, + "semantic": 0.5, + "toin": 0.5 + }, + "semantic_score": 0.5, + "toin_score": 0.5, + "tokens": 0, + "total_score": 0.47954 + }, + { + "density_score": 1.0, + "drop_safe": true, + "error_score": 0.0, + "is_protected": false, + "message_index": 2, + "recency_score": 0.90484, + "reference_score": 0.0, + "score_breakdown": { + "density": 1.0, + "error": 0.0, + "recency": 0.90484, + "reference": 0.0, + "semantic": 0.5, + "toin": 0.5 + }, + "semantic_score": 0.5, + "toin_score": 0.5, + "tokens": 8, + "total_score": 0.45597 + }, + { + "density_score": 1.0, + "drop_safe": true, + "error_score": 0.0, + "is_protected": false, + "message_index": 3, + "recency_score": 1.0, + "reference_score": 0.0, + "score_breakdown": { + "density": 1.0, + "error": 0.0, + "recency": 1.0, + "reference": 0.0, + "semantic": 0.5, + "toin": 0.5 + }, + "semantic_score": 0.5, + "toin_score": 0.5, + "tokens": 4, + "total_score": 0.475 + } + ], + "recorded_at": "2026-05-01T04:29:36.074466+00:00", + "transform": "message_scorer" +} diff --git a/tests/parity/fixtures/message_scorer/unicode_content_22231f9bffa6.json b/tests/parity/fixtures/message_scorer/unicode_content_22231f9bffa6.json new file mode 100644 index 000000000..686ff63cc --- /dev/null +++ b/tests/parity/fixtures/message_scorer/unicode_content_22231f9bffa6.json @@ -0,0 +1,68 @@ +{ + "config": { + "weights": null + }, + "input": { + "decay_rate": 0.1, + "messages": [ + { + "content": "\u4f60\u597d\u4e16\u754c \u044d\u0442\u043e \u0442\u0435\u0441\u0442 \ud83c\udf89\ud83c\udf89\ud83c\udf89", + "role": "user" + }, + { + "content": "received the unicode message", + "role": "assistant" + } + ], + "protected_indices": [], + "tool_unit_indices": [] + }, + "input_sha256": "22231f9bffa6dadeb455f2d1515019d7924ed307fe27331b2fb8720297b0e2b4", + "label": "unicode_content", + "output": [ + { + "density_score": 1.0, + "drop_safe": true, + "error_score": 0.0, + "is_protected": false, + "message_index": 0, + "recency_score": 0.90484, + "reference_score": 0.0, + "score_breakdown": { + "density": 1.0, + "error": 0.0, + "recency": 0.90484, + "reference": 0.0, + "semantic": 0.5, + "toin": 0.5 + }, + "semantic_score": 0.5, + "toin_score": 0.5, + "tokens": 4, + "total_score": 0.45597 + }, + { + "density_score": 1.0, + "drop_safe": true, + "error_score": 0.0, + "is_protected": false, + "message_index": 1, + "recency_score": 1.0, + "reference_score": 0.0, + "score_breakdown": { + "density": 1.0, + "error": 0.0, + "recency": 1.0, + "reference": 0.0, + "semantic": 0.5, + "toin": 0.5 + }, + "semantic_score": 0.5, + "toin_score": 0.5, + "tokens": 7, + "total_score": 0.475 + } + ], + "recorded_at": "2026-05-01T04:29:36.076815+00:00", + "transform": "message_scorer" +} diff --git a/tests/parity/record_message_scorer.py b/tests/parity/record_message_scorer.py new file mode 100644 index 000000000..79a8bf896 --- /dev/null +++ b/tests/parity/record_message_scorer.py @@ -0,0 +1,350 @@ +"""Record `MessageScorer` parity fixtures. + +Captures `MessageScorer.score_messages(messages, protected, tool_unit)` +with `toin=None` and `embedding_provider=None` so all six factors run +through the deterministic-or-neutral code path. The Rust comparator +(`MessageScorerComparator` in `crates/headroom-parity/src/lib.rs`) runs +the same inputs through the Rust port and asserts bit-equal outputs +after a 5-decimal-place rounding step. + +Why round: Rust uses `f32::exp` and computes the weighted total in +f32, while Python's `math.exp` and weighted sum are f64. Both are +mathematically identical; they only drift in the low bits. Rounding +both sides to 5 decimals (1e-5 tolerance, 100x looser than f32 ulp +drift on the 0–1 score range) gives byte-equal JSON without masking +real bugs. + +Run from repo root: + python tests/parity/record_message_scorer.py +""" + +from __future__ import annotations + +import datetime as _dt +import hashlib +import json +from dataclasses import asdict +from pathlib import Path +from typing import Any + +from headroom.config import ScoringWeights +from headroom.transforms.scoring import MessageScorer + +_REPO_ROOT = Path(__file__).resolve().parent.parent.parent +_FIXTURES_DIR = _REPO_ROOT / "tests" / "parity" / "fixtures" / "message_scorer" + +# 5 decimals: f32 has ~7 decimals of precision, but Rust's port runs +# the weighted-sum in f32 while Python runs it in f64 — six summed +# f32-precision components occasionally drift in the 6th decimal of +# the total. 5 decimals (1e-5 tolerance) is loose enough to absorb +# the drift while still tight enough to catch real bugs. +_FLOAT_ROUND_PLACES = 5 + + +def _round_floats(obj: Any) -> Any: + """Recursively round every float in a JSON-shaped object.""" + if isinstance(obj, float): + return round(obj, _FLOAT_ROUND_PLACES) + if isinstance(obj, dict): + return {k: _round_floats(v) for k, v in obj.items()} + if isinstance(obj, list): + return [_round_floats(v) for v in obj] + return obj + + +def _digest(payload: dict[str, Any]) -> str: + blob = json.dumps(payload, sort_keys=True).encode("utf-8") + return hashlib.sha256(blob).hexdigest() + + +def _record( + label: str, + messages: list[dict[str, Any]], + protected_indices: list[int], + tool_unit_indices: list[int], + weights: ScoringWeights | None = None, + decay_rate: float = 0.1, +) -> Path: + scorer = MessageScorer( + weights=weights, + toin=None, + embedding_provider=None, + recency_decay_rate=decay_rate, + ) + scores = scorer.score_messages( + messages=messages, + protected_indices=set(protected_indices), + tool_unit_indices=set(tool_unit_indices), + ) + + payload_input = { + "messages": messages, + "protected_indices": sorted(protected_indices), + "tool_unit_indices": sorted(tool_unit_indices), + "decay_rate": decay_rate, + } + payload_config = {"weights": asdict(weights) if weights else None} + payload_output = _round_floats([asdict(s) for s in scores]) + + digest_source = { + "transform": "message_scorer", + "label": label, + "input": payload_input, + "config": payload_config, + } + digest = _digest(digest_source) + + fixture = { + "transform": "message_scorer", + "label": label, + "input": payload_input, + "config": payload_config, + "output": payload_output, + "recorded_at": _dt.datetime.now(tz=_dt.timezone.utc).isoformat(), + "input_sha256": digest, + } + + _FIXTURES_DIR.mkdir(parents=True, exist_ok=True) + target = _FIXTURES_DIR / f"{label}_{digest[:12]}.json" + target.write_text(json.dumps(fixture, indent=2, sort_keys=True) + "\n") + return target + + +def _scenarios() -> list[dict[str, Any]]: + """Test scenarios covering each deterministic factor + edge cases.""" + out: list[dict[str, Any]] = [] + + # 1. Single message — recency=1.0, no refs, no tool unit. + out.append( + { + "label": "single_user_message", + "messages": [{"role": "user", "content": "hello world"}], + "protected_indices": [], + "tool_unit_indices": [], + } + ) + + # 2. Empty list. + out.append( + { + "label": "empty", + "messages": [], + "protected_indices": [], + "tool_unit_indices": [], + } + ) + + # 3. Linear conversation (5 messages, no tools) — exercises recency + # decay over a small range. + out.append( + { + "label": "linear_5_messages", + "messages": [ + {"role": "user", "content": "what is the capital of france"}, + {"role": "assistant", "content": "the capital of france is paris"}, + {"role": "user", "content": "and germany"}, + {"role": "assistant", "content": "the capital of germany is berlin"}, + {"role": "user", "content": "thanks"}, + ], + "protected_indices": [0], + "tool_unit_indices": [], + } + ) + + # 4. Tool-call pair — exercises forward references. + out.append( + { + "label": "tool_call_pair", + "messages": [ + {"role": "user", "content": "what's the weather"}, + { + "role": "assistant", + "content": "", + "tool_calls": [ + { + "id": "call_1", + "type": "function", + "function": {"name": "get_weather", "arguments": "{}"}, + } + ], + }, + { + "role": "tool", + "tool_call_id": "call_1", + "content": '{"temp": 72, "conditions": "sunny"}', + }, + {"role": "assistant", "content": "it is 72 and sunny"}, + ], + "protected_indices": [0], + "tool_unit_indices": [1, 2], + } + ) + + # 5. Multiple tool-call references to same assistant message. + out.append( + { + "label": "multi_ref_to_assistant", + "messages": [ + {"role": "user", "content": "do many things"}, + { + "role": "assistant", + "content": "", + "tool_calls": [ + {"id": "c1", "function": {"name": "f"}}, + {"id": "c2", "function": {"name": "g"}}, + {"id": "c3", "function": {"name": "h"}}, + ], + }, + {"role": "tool", "tool_call_id": "c1", "content": "r1"}, + {"role": "tool", "tool_call_id": "c2", "content": "r2"}, + {"role": "tool", "tool_call_id": "c3", "content": "r3"}, + ], + "protected_indices": [], + "tool_unit_indices": [1, 2, 3, 4], + } + ) + + # 6. High-density message (all-unique tokens). + out.append( + { + "label": "high_density", + "messages": [ + {"role": "user", "content": "alpha bravo charlie delta echo foxtrot golf hotel"}, + ], + "protected_indices": [], + "tool_unit_indices": [], + } + ) + + # 7. Low-density (highly repetitive). + out.append( + { + "label": "low_density", + "messages": [ + {"role": "user", "content": "ok ok ok ok ok ok ok ok ok ok"}, + ], + "protected_indices": [], + "tool_unit_indices": [], + } + ) + + # 8. Custom weights — exercises ScoringWeights normalization + + # weighted total. + out.append( + { + "label": "custom_weights_recency_heavy", + "messages": [ + {"role": "user", "content": "first message in a longer chat"}, + {"role": "assistant", "content": "an assistant reply with substance"}, + {"role": "user", "content": "another follow up question here"}, + {"role": "assistant", "content": "and the closing reply"}, + ], + "protected_indices": [], + "tool_unit_indices": [], + "weights": ScoringWeights( + recency=0.6, + semantic_similarity=0.1, + toin_importance=0.1, + error_indicator=0.05, + forward_reference=0.1, + token_density=0.05, + ), + } + ) + + # 9. Custom decay rate (slower decay). + out.append( + { + "label": "slow_decay", + "messages": [{"role": "user", "content": f"message number {i}"} for i in range(8)], + "protected_indices": [], + "tool_unit_indices": [], + "decay_rate": 0.02, + } + ) + + # 10. drop_safe truth table — protected + in_tool_unit. + out.append( + { + "label": "drop_safe_truth_table", + "messages": [ + {"role": "system", "content": "you are helpful"}, + {"role": "user", "content": "hello"}, + { + "role": "assistant", + "content": "", + "tool_calls": [{"id": "x", "function": {"name": "f"}}], + }, + {"role": "tool", "tool_call_id": "x", "content": "result"}, + ], + "protected_indices": [0, 2], + "tool_unit_indices": [2, 3], + } + ) + + # 11. Unicode content — exercises char-count tokens estimate. + out.append( + { + "label": "unicode_content", + "messages": [ + {"role": "user", "content": "你好世界 это тест 🎉🎉🎉"}, + {"role": "assistant", "content": "received the unicode message"}, + ], + "protected_indices": [], + "tool_unit_indices": [], + } + ) + + # 12. Tool-call id mismatch — tool message references a non-existent + # call_id, must NOT contribute to forward refs. + out.append( + { + "label": "orphan_tool_response", + "messages": [ + {"role": "user", "content": "go"}, + {"role": "tool", "tool_call_id": "nope", "content": "stranded"}, + ], + "protected_indices": [], + "tool_unit_indices": [], + } + ) + + # 13. Non-string content (tool_calls list, no text content). + out.append( + { + "label": "non_string_content", + "messages": [ + {"role": "user", "content": "hi"}, + { + "role": "assistant", + "content": None, + "tool_calls": [{"id": "tc", "function": {"name": "f"}}], + }, + ], + "protected_indices": [], + "tool_unit_indices": [], + } + ) + + return out + + +def main() -> int: + written: list[Path] = [] + for sc in _scenarios(): + path = _record( + label=sc["label"], + messages=sc["messages"], + protected_indices=sc["protected_indices"], + tool_unit_indices=sc["tool_unit_indices"], + weights=sc.get("weights"), + decay_rate=sc.get("decay_rate", 0.1), + ) + written.append(path) + print(f" + {path.relative_to(_REPO_ROOT)}") + print(f"wrote {len(written)} fixture(s) → {_FIXTURES_DIR.relative_to(_REPO_ROOT)}") + return 0 + + +if __name__ == "__main__": + raise SystemExit(main())