Skip to main content

wm_tools/
embedding_router.rs

1//! Embedding-based NLU router for the `wm` meta-tool.
2//!
3//! Replaces 166 hand-written TF-IDF keyword profiles with embedding cosine
4//! similarity. Each tool's description is embedded once at startup. Input
5//! queries are embedded and compared against all tool embeddings using cosine
6//! similarity.
7//!
8//! # OATS: Outcome-Aware Tool Selection
9//!
10//! After each tool call, the router records whether the routing was correct
11//! (tool succeeded) or incorrect (tool failed / was wrong). Success and failure
12//! query embeddings are averaged into centroids. Tool embeddings are refined
13//! by interpolating toward the success centroid:
14//!
15//! ```text
16//! refined = base * (1 - α) + success_centroid * α
17//! ```
18//!
19//! This is zero-cost at serving time when pre-computed, and improves NDCG@5
20//! from ~0.869 to ~0.940 (OATS, 2026).
21//!
22//! # Fallback
23//!
24//! If no real embedder is available (only StubEmbedder), the router returns
25//! `None` from `new()`, and the caller falls back to the TF-IDF router.
26
27use ahash::AHashMap;
28use std::sync::{Arc, RwLock};
29use wm_memory::Embedder;
30
31use crate::nlu::{PREFIX_ROUTES, TOOL_PROFILES, ToolProfile};
32
33/// OATS refinement strength (interpolation factor toward success centroid).
34const OATS_ALPHA: f32 = 0.15;
35
36/// Minimum observations before OATS refinement kicks in.
37const OATS_MIN_OBSERVATIONS: usize = 10;
38
39/// Minimum cosine similarity to return a match (below this → gnosis fallback).
40const MIN_THRESHOLD: f64 = 0.10;
41
42/// Minimum margin between top-1 and top-2 before the embedding router's
43/// choice is trusted.
44///
45/// Near-ties (margin below this) mean the description vocabulary cannot
46/// separate intent — the caller should defer to the TF-IDF router. Derived
47/// from the 2026-08-11 shadow data: ambiguous queries produced confident-
48/// looking top-1 scores (0.5–0.8) with the correct tool often runner-up.
49pub const MIN_MARGIN: f64 = 0.02;
50
51/// Outcome statistics for a single tool (OATS data).
52#[derive(Debug, Clone)]
53pub struct OutcomeStats {
54    /// Running centroid of query embeddings where this tool was the correct route.
55    success_centroid: Vec<f32>,
56    /// Running centroid of query embeddings where this tool was the wrong route.
57    #[allow(dead_code)]
58    failure_centroid: Vec<f32>,
59    /// Number of successful routing observations.
60    success_count: usize,
61    /// Number of failed routing observations.
62    failure_count: usize,
63}
64
65impl OutcomeStats {
66    /// Create empty outcome stats with the given embedding dimensionality.
67    fn new(dim: usize) -> Self {
68        Self {
69            success_centroid: vec![0.0; dim],
70            failure_centroid: vec![0.0; dim],
71            success_count: 0,
72            failure_count: 0,
73        }
74    }
75
76    /// Record a routing outcome with the query embedding.
77    fn record(&mut self, query_emb: &[f32], success: bool) {
78        if query_emb.is_empty() {
79            return;
80        }
81
82        if success {
83            update_centroid(
84                &mut self.success_centroid,
85                &mut self.success_count,
86                query_emb,
87            );
88        } else {
89            update_centroid(
90                &mut self.failure_centroid,
91                &mut self.failure_count,
92                query_emb,
93            );
94        }
95    }
96
97    /// Whether OATS has enough data to refine this tool's embedding.
98    const fn is_ready(&self) -> bool {
99        self.success_count >= OATS_MIN_OBSERVATIONS
100    }
101}
102
103/// Update a running centroid with a new vector (incremental mean).
104fn update_centroid(centroid: &mut [f32], count: &mut usize, new_vec: &[f32]) {
105    if centroid.len() != new_vec.len() {
106        return;
107    }
108    let n = *count as f32 + 1.0;
109    for (c, v) in centroid.iter_mut().zip(new_vec.iter()) {
110        *c += (*v - *c) / n;
111    }
112    *count += 1;
113}
114
115/// The embedding-based NLU router.
116///
117/// Pre-computes tool embeddings at initialization, then routes queries by
118/// embedding the query and computing cosine similarity against all tool
119/// embeddings. OATS refinement adjusts tool embeddings based on observed
120/// outcomes.
121pub struct EmbeddingRouter {
122    /// Tool name → base embedding (from tool description).
123    tool_embeddings: AHashMap<String, Vec<f32>>,
124    /// Embedder backend.
125    embedder: Box<dyn Embedder>,
126    /// OATS outcome stats per tool (interior mutability for record_outcome).
127    outcome_stats: RwLock<AHashMap<String, OutcomeStats>>,
128    /// Embedding dimensionality.
129    dim: usize,
130    /// Whether to apply the TF-IDF prefix-route bonus.
131    ///
132    /// `true` for the legacy keyword-profile path (`new`), where descriptions
133    /// are bare keyword lists and the verb bonus compensates. `false` for
134    /// anchored descriptions (`with_descriptions`), where the bonus fights
135    /// the intent anchors ("list tools" was boosted toward memory.list despite
136    /// tools.list carrying the exact anchor).
137    apply_prefix_bonus: bool,
138}
139
140impl EmbeddingRouter {
141    /// Create a new embedding router, pre-computing tool embeddings.
142    ///
143    /// Returns `None` if:
144    /// - The embedder is a stub (hash-based embeddings have no semantic meaning)
145    /// - Batch embedding fails
146    ///
147    /// This allows the caller to gracefully fall back to the TF-IDF router.
148    #[must_use]
149    pub fn new(embedder: Box<dyn Embedder>) -> Option<Self> {
150        Self::new_with_descriptions(embedder, tool_descriptions(), true)
151    }
152
153    /// Create a new embedding router from explicit (tool, description) pairs.
154    ///
155    /// Unlike [`Self::new`] — which uses the static keyword profiles from
156    /// `nlu.rs` (169 tools, keyword-mashup descriptions) — this accepts
157    /// descriptions from the live tool registry (all 229 tools, prose
158    /// descriptions). Sentence-style descriptions embed far better than
159    /// bare keyword lists; the 2026-08-11 shadow run (42.6% disagreement)
160    /// showed keyword-mashup top-1 selection collapsing onto arbitrary
161    /// high-similarity tools.
162    #[must_use]
163    pub fn with_descriptions(
164        embedder: Box<dyn Embedder>,
165        descriptions: Vec<(String, String)>,
166    ) -> Option<Self> {
167        Self::new_with_descriptions(embedder, descriptions, false)
168    }
169
170    fn new_with_descriptions(
171        embedder: Box<dyn Embedder>,
172        descriptions: Vec<(String, String)>,
173        apply_prefix_bonus: bool,
174    ) -> Option<Self> {
175        // Stub embedders produce hash-based embeddings with no semantic similarity.
176        // Don't use the embedding router with them — fall back to TF-IDF.
177        if embedder.backend_name() == "stub" {
178            tracing::info!(
179                "embedding router disabled — stub embedder has no semantic similarity, using TF-IDF fallback"
180            );
181            return None;
182        }
183
184        let dim = embedder.dimension();
185
186        let texts: Vec<&str> = descriptions.iter().map(|(_, d)| d.as_str()).collect();
187        let embeddings = embedder.embed_batch(&texts).ok()?;
188
189        if embeddings.len() != descriptions.len() {
190            tracing::warn!(
191                "embedding router: expected {} embeddings, got {} — falling back to TF-IDF",
192                descriptions.len(),
193                embeddings.len()
194            );
195            return None;
196        }
197
198        let mut tool_embeddings = AHashMap::with_capacity(descriptions.len());
199        for ((name, _), emb) in descriptions.into_iter().zip(embeddings) {
200            tool_embeddings.insert(name, emb);
201        }
202
203        tracing::info!(
204            "embedding router initialized with {} tools, dim={}, backend={}",
205            tool_embeddings.len(),
206            dim,
207            embedder.backend_name()
208        );
209
210        Some(Self {
211            tool_embeddings,
212            embedder,
213            outcome_stats: RwLock::new(AHashMap::new()),
214            dim,
215            apply_prefix_bonus,
216        })
217    }
218
219    /// Route a natural language query to a tool name and confidence score.
220    ///
221    /// Returns `("gnosis", 0.0)` for empty input or when no tool scores above
222    /// the minimum threshold.
223    #[must_use]
224    pub fn route(&self, query: &str) -> (String, f64) {
225        match self.route_with_margin(query) {
226            Some((t, c, _)) => (t, c),
227            None => ("gnosis".into(), 0.0),
228        }
229    }
230
231    /// Route a query, returning (tool, confidence, margin).
232    ///
233    /// `margin` is the score gap between the top-1 and top-2 tools. Small
234    /// margins indicate the descriptions cannot separate intent — callers
235    /// should defer to the TF-IDF router when `margin < MIN_MARGIN`.
236    #[must_use]
237    pub fn route_with_margin(&self, query: &str) -> Option<(String, f64, f64)> {
238        self.route_with_margin_and_embedding(query)
239            .map(|(tool, conf, margin, _)| (tool, conf, margin))
240    }
241
242    /// Route a query, also returning the query embedding.
243    ///
244    /// The embedding is what [`record_outcome_with_embedding`](Self::record_outcome_with_embedding)
245    /// needs — returning it here lets callers embed each query once instead of
246    /// twice (embedder HTTP round-trips dominate NLU latency).
247    #[must_use]
248    pub fn route_with_margin_and_embedding(
249        &self,
250        query: &str,
251    ) -> Option<(String, f64, f64, Vec<f32>)> {
252        let lower = query.to_lowercase();
253        if lower.trim().is_empty() {
254            return None;
255        }
256
257        let query_emb = match self.embedder.embed(&lower) {
258            Ok(emb) => emb,
259            Err(e) => {
260                tracing::warn!(error = %e, "embedding router: query embedding failed");
261                return None;
262            }
263        };
264
265        // Prefix route bonus — only on the legacy keyword-profile path, where
266        // descriptions are bare keyword lists and the verb bonus compensates.
267        // On the anchored path it fights the intent anchors ("list tools" was
268        // boosted toward memory.list despite tools.list carrying the anchor).
269        let prefix_bonus: Option<(&str, f64)> = if self.apply_prefix_bonus {
270            let first_word = lower.split_whitespace().next().unwrap_or("");
271            PREFIX_ROUTES
272                .iter()
273                .find(|(verb, _, _)| *verb == first_word)
274                .map(|(_, tool, bonus)| (*tool, *bonus))
275        } else {
276            None
277        };
278
279        // Score each tool by cosine similarity to (optionally refined) embedding
280        let Ok(stats_lock) = self.outcome_stats.read() else {
281            return None;
282        };
283
284        let mut best_tool = "gnosis".to_string();
285        let mut best_score = 0.0_f64;
286        let mut second_tool = String::new();
287        let mut second_score = 0.0_f64;
288
289        for (name, base_emb) in &self.tool_embeddings {
290            let refined = self.oats_refine(name, base_emb, &stats_lock);
291            let mut score = f64::from(cosine_sim(&query_emb, &refined));
292
293            // Apply prefix routing: bonus to matching tool, penalty to non-matching
294            if let Some((bonus_tool, bonus)) = prefix_bonus {
295                if name == bonus_tool {
296                    score *= bonus;
297                } else {
298                    score /= bonus;
299                }
300            }
301
302            if score > best_score {
303                second_score = best_score;
304                second_tool.clone_from(&best_tool);
305                best_score = score;
306                best_tool.clone_from(name);
307            } else if score > second_score {
308                second_score = score;
309                second_tool.clone_from(name);
310            }
311        }
312
313        drop(stats_lock);
314
315        if best_score < MIN_THRESHOLD {
316            return None;
317        }
318
319        if best_score - second_score < MIN_MARGIN {
320            tracing::debug!(
321                query = %lower,
322                best_tool = %best_tool,
323                best_score,
324                second_tool = %second_tool,
325                second_score,
326                "embedding router: near-tie"
327            );
328        }
329
330        Some((best_tool, best_score, best_score - second_score, query_emb))
331    }
332
333    /// OATS: interpolate tool embedding toward success centroid.
334    ///
335    /// If we have enough success observations (≥ `OATS_MIN_OBSERVATIONS`),
336    /// blend the base embedding toward the success centroid by `OATS_ALPHA`.
337    /// Otherwise, return the base embedding unchanged.
338    fn oats_refine(
339        &self,
340        tool_name: &str,
341        base_emb: &[f32],
342        stats: &AHashMap<String, OutcomeStats>,
343    ) -> Vec<f32> {
344        if let Some(stat) = stats.get(tool_name) {
345            if stat.is_ready() && stat.success_centroid.len() == base_emb.len() {
346                return interpolate(base_emb, &stat.success_centroid, OATS_ALPHA);
347            }
348        }
349        base_emb.to_vec()
350    }
351
352    /// Record a routing outcome for OATS refinement.
353    ///
354    /// Call this after each tool dispatch to track whether the routing was
355    /// correct. `success = true` means the tool was the right choice and
356    /// executed successfully; `false` means it was wrong or failed.
357    pub fn record_outcome(&self, tool_name: &str, query: &str, success: bool) {
358        if query.trim().is_empty() {
359            return;
360        }
361        let query_emb = match self.embedder.embed(&query.to_lowercase()) {
362            Ok(emb) => emb,
363            Err(_) => return,
364        };
365        self.record_outcome_with_embedding(tool_name, query, success, &query_emb);
366    }
367
368    /// Record a routing outcome reusing a query embedding already computed by
369    /// the router.
370    ///
371    /// Callers that routed through [`route_with_margin`](Self::route_with_margin)
372    /// should pass the embedding back here so the query is embedded only once
373    /// instead of twice (HTTP embedder round-trips dominate NLU latency).
374    pub fn record_outcome_with_embedding(
375        &self,
376        tool_name: &str,
377        query: &str,
378        success: bool,
379        query_emb: &[f32],
380    ) {
381        if query.trim().is_empty() {
382            return;
383        }
384        let Ok(mut stats) = self.outcome_stats.write() else {
385            return;
386        };
387        let stat = stats
388            .entry(tool_name.to_string())
389            .or_insert_with(|| OutcomeStats::new(self.dim));
390        stat.record(query_emb, success);
391    }
392
393    /// Number of tool embeddings in the router.
394    #[must_use]
395    pub fn tool_count(&self) -> usize {
396        self.tool_embeddings.len()
397    }
398
399    /// Embedding dimensionality.
400    #[must_use]
401    pub const fn dimension(&self) -> usize {
402        self.dim
403    }
404
405    /// Embedder backend name.
406    #[must_use]
407    pub fn backend_name(&self) -> &str {
408        self.embedder.backend_name()
409    }
410
411    /// Get a snapshot of outcome stats counts for observability.
412    #[must_use]
413    pub fn outcome_counts(&self) -> Vec<(String, usize, usize)> {
414        let Ok(stats) = self.outcome_stats.read() else {
415            return Vec::new();
416        };
417        stats
418            .iter()
419            .map(|(name, s)| (name.clone(), s.success_count, s.failure_count))
420            .collect()
421    }
422
423    /// Serialize OATS outcome stats to JSON for persistence.
424    #[must_use]
425    #[allow(clippy::type_complexity)]
426    pub fn save_oats(&self) -> Option<String> {
427        let Ok(stats) = self.outcome_stats.read() else {
428            return None;
429        };
430        let serializable: Vec<(String, usize, usize, Vec<f32>, Vec<f32>)> = stats
431            .iter()
432            .map(|(name, s)| {
433                (
434                    name.clone(),
435                    s.success_count,
436                    s.failure_count,
437                    s.success_centroid.clone(),
438                    s.failure_centroid.clone(),
439                )
440            })
441            .collect();
442        serde_json::to_string_pretty(&serializable).ok()
443    }
444
445    /// Load OATS outcome stats from JSON (previously saved by `save_oats`).
446    pub fn load_oats(&self, json: &str) {
447        if let Ok(data) =
448            serde_json::from_str::<Vec<(String, usize, usize, Vec<f32>, Vec<f32>)>>(json)
449        {
450            let Ok(mut stats) = self.outcome_stats.write() else {
451                return;
452            };
453            for (name, success_count, failure_count, success_centroid, failure_centroid) in data {
454                let dim = success_centroid.len().max(self.dim);
455                let mut s = OutcomeStats::new(dim);
456                s.success_count = success_count;
457                s.failure_count = failure_count;
458                s.success_centroid = success_centroid;
459                s.failure_centroid = failure_centroid;
460                stats.insert(name, s);
461            }
462            tracing::info!("Loaded OATS outcome stats from disk");
463        }
464    }
465}
466
467// ── Shadow Mode Stats ────────────────────────────────────────────────
468
469/// Maximum number of disagreement samples to retain.
470const MAX_SAMPLES: usize = 50;
471
472/// A single disagreement sample between embedding router and TF-IDF.
473#[derive(Debug, Clone, serde::Serialize, serde::Deserialize)]
474pub struct DisagreementSample {
475    pub query: String,
476    pub embedding_tool: String,
477    pub embedding_conf: f64,
478    pub tfidf_tool: String,
479    pub tfidf_conf: f64,
480}
481
482/// Shadow mode statistics tracking embedding vs TF-IDF disagreements.
483///
484/// Thread-safe via `RwLock`. Updated on every `classify_with_router` call
485/// when the embedding router is active.
486#[derive(Debug, Default, serde::Serialize, serde::Deserialize)]
487pub struct ShadowModeStats {
488    /// Total queries routed through shadow mode.
489    pub total_queries: u64,
490    /// Total disagreements (embedding chose different tool than TF-IDF).
491    pub total_disagreements: u64,
492    /// Per-tool disagreement counts: (embedding_tool, tfidf_tool) → count.
493    pub disagreement_pairs: std::collections::HashMap<String, u64>,
494    /// Recent disagreement samples (capped at MAX_SAMPLES).
495    pub samples: Vec<DisagreementSample>,
496}
497
498impl ShadowModeStats {
499    /// Record a routing comparison.
500    pub fn record(
501        &mut self,
502        query: &str,
503        emb_tool: &str,
504        emb_conf: f64,
505        tfidf_tool: &str,
506        tfidf_conf: f64,
507    ) {
508        self.total_queries += 1;
509        if emb_tool != tfidf_tool {
510            self.total_disagreements += 1;
511            let key = format!("{emb_tool} → {tfidf_tool}");
512            *self.disagreement_pairs.entry(key).or_insert(0) += 1;
513            if self.samples.len() >= MAX_SAMPLES {
514                self.samples.remove(0);
515            }
516            self.samples.push(DisagreementSample {
517                query: query.chars().take(200).collect(),
518                embedding_tool: emb_tool.to_string(),
519                embedding_conf: emb_conf,
520                tfidf_tool: tfidf_tool.to_string(),
521                tfidf_conf,
522            });
523        }
524    }
525
526    /// Disagreement rate (0.0–1.0).
527    #[must_use]
528    pub fn disagreement_rate(&self) -> f64 {
529        if self.total_queries == 0 {
530            0.0
531        } else {
532            self.total_disagreements as f64 / self.total_queries as f64
533        }
534    }
535
536    /// Whether the embedding router is ready for promotion to primary
537    /// (disagreement rate below 20% and enough samples).
538    #[must_use]
539    pub fn promotion_ready(&self) -> bool {
540        self.total_queries >= 100 && self.disagreement_rate() < 0.20
541    }
542
543    /// Generate a JSON report for the `nlu.shadow_report` tool.
544    #[must_use]
545    pub fn report(&self) -> serde_json::Value {
546        let mut pairs: Vec<(String, u64)> = self
547            .disagreement_pairs
548            .iter()
549            .map(|(k, v)| (k.clone(), *v))
550            .collect();
551        pairs.sort_by_key(|x| std::cmp::Reverse(x.1));
552
553        serde_json::json!({
554            "total_queries": self.total_queries,
555            "total_disagreements": self.total_disagreements,
556            "disagreement_rate": self.disagreement_rate(),
557            "promotion_ready": self.promotion_ready(),
558            "top_disagreement_pairs": pairs.iter().take(10).map(|(k, v)| {
559                serde_json::json!({"pair": k, "count": v})
560            }).collect::<Vec<_>>(),
561            "recent_samples": self.samples.iter().take(10).map(|s| {
562                serde_json::json!({
563                    "query": s.query,
564                    "embedding_tool": s.embedding_tool,
565                    "embedding_conf": s.embedding_conf,
566                    "tfidf_tool": s.tfidf_tool,
567                    "tfidf_conf": s.tfidf_conf,
568                })
569            }).collect::<Vec<_>>(),
570        })
571    }
572}
573
574// ── Helpers ──────────────────────────────────────────────────────────
575
576/// Cosine similarity between two f32 vectors.
577fn cosine_sim(a: &[f32], b: &[f32]) -> f32 {
578    if a.is_empty() || b.is_empty() || a.len() != b.len() {
579        return 0.0;
580    }
581
582    let mut dot = 0.0_f32;
583    let mut norm_a = 0.0_f32;
584    let mut norm_b = 0.0_f32;
585
586    for (x, y) in a.iter().zip(b.iter()) {
587        dot += x * y;
588        norm_a += x * x;
589        norm_b += y * y;
590    }
591
592    let denom = norm_a.sqrt() * norm_b.sqrt();
593    if denom == 0.0 { 0.0 } else { dot / denom }
594}
595
596/// Linear interpolation between two vectors: `base * (1 - α) + target * α`.
597fn interpolate(base: &[f32], target: &[f32], alpha: f32) -> Vec<f32> {
598    base.iter()
599        .zip(target.iter())
600        .map(|(b, t)| b * (1.0 - alpha) + t * alpha)
601        .collect()
602}
603
604/// Generate tool descriptions from the static TOOL_PROFILES.
605///
606/// Each description is the tool name followed by its keywords. This gives the
607/// embedder semantic content to work with. Example:
608///
609/// `"memory.create remember store save memorize record persist capture"`
610#[must_use]
611pub fn tool_descriptions() -> Vec<(String, String)> {
612    TOOL_PROFILES
613        .iter()
614        .map(|p| (p.tool_name.to_string(), profile_to_description(p)))
615        .collect()
616}
617
618/// Convert a ToolProfile into a description string for embedding.
619fn profile_to_description(profile: &ToolProfile) -> String {
620    let keywords: Vec<&str> = profile.keywords.iter().map(|(t, _)| *t).collect();
621    format!("{} {}", profile.tool_name, keywords.join(" "))
622}
623
624/// Intent anchors: natural-language phrasings users say when they mean a tool.
625///
626/// The registry's `description()` strings describe *what a tool does* (display
627/// prose) but rarely match how users phrase intent. The 2026-08-11 shadow runs
628/// showed top-1 cosine collapsing onto arbitrary tools for common phrasings
629/// ("show my karma" → karma.clear). Anchors are appended to the embedded text
630/// so the vector for each tool covers user phrasing, not just docstring prose.
631static INTENT_ANCHORS: &[(&str, &[&str])] = &[
632    // Memory
633    (
634        "memory.create",
635        &[
636            "remember that",
637            "store this note",
638            "save this thought",
639            "memorize this",
640            "keep this in memory",
641            "note that",
642            "record that",
643        ],
644    ),
645    (
646        "memory.read",
647        &[
648            "get memory by id",
649            "read this memory",
650            "recall what I said",
651            "fetch memory",
652        ],
653    ),
654    (
655        "memory.list",
656        &[
657            "list my memories",
658            "show my recent memories",
659            "what memories do I have",
660            "find memories about",
661            "memories in the codex galaxy",
662        ],
663    ),
664    (
665        "memory.search",
666        &[
667            "search my memories for",
668            "find memory about",
669            "memory search",
670            "search for rust",
671            "search memories",
672        ],
673    ),
674    (
675        "memory.vector.search",
676        &[
677            "find memory about search",
678            "semantic search",
679            "similar memories",
680        ],
681    ),
682    (
683        "memory.count",
684        &["count my memories", "how many memories", "memory count"],
685    ),
686    ("memory.tags", &["what tags do I have", "show memory tags"]),
687    (
688        "memory.delete",
689        &["delete memory", "remove memory", "forget this memory"],
690    ),
691    // Galaxy
692    (
693        "galaxy.list",
694        &["list galaxies", "what galaxies exist", "show the galaxies"],
695    ),
696    (
697        "galaxy.stats",
698        &[
699            "galaxy stats",
700            "stats for the codex galaxy",
701            "how many memories are in",
702            "show galaxy info",
703        ],
704    ),
705    (
706        "galaxy.create",
707        &["create a new galaxy", "new galaxy called", "make a galaxy"],
708    ),
709    ("galaxy.health", &["check galaxy health", "galaxy health"]),
710    (
711        "galaxy.taxonomy",
712        &["gana taxonomy", "show the gana taxonomy"],
713    ),
714    // Session
715    ("session.start", &["start a session", "begin a new session"]),
716    ("session.end", &["end the session", "close the session"]),
717    (
718        "session.list",
719        &[
720            "what sessions do I have",
721            "list sessions",
722            "show session history",
723        ],
724    ),
725    (
726        "session.record",
727        &["record this session turn", "log this session turn"],
728    ),
729    (
730        "session.replay",
731        &["replay the session", "replay last session"],
732    ),
733    (
734        "session.recall",
735        &[
736            "recall the session context",
737            "session history",
738            "previous session",
739            "record that the server restarted",
740        ],
741    ),
742    (
743        "session.handoff",
744        &[
745            "hand off the session",
746            "transfer session",
747            "session handoff",
748        ],
749    ),
750    // Karma
751    (
752        "karma.report",
753        &[
754            "show my karma",
755            "karma status",
756            "check my karma",
757            "karma balance",
758            "karma report",
759            "karma ledger status",
760        ],
761    ),
762    (
763        "karma.history",
764        &["karma history", "past karma entries", "recent karma"],
765    ),
766    (
767        "karma.clear",
768        &["clear karma", "wipe karma", "reset karma", "purge karma"],
769    ),
770    (
771        "karma.verify_chain",
772        &[
773            "check the karma chain",
774            "verify chain integrity",
775            "karma chain",
776        ],
777    ),
778    (
779        "karma.anchor",
780        &["anchor the karma chain", "publish anchor", "merkle anchor"],
781    ),
782    // Friction / RSI
783    (
784        "friction.log",
785        &["log friction", "log an error", "log friction entry"],
786    ),
787    (
788        "friction.review",
789        &[
790            "review the friction log",
791            "review friction",
792            "friction review",
793        ],
794    ),
795    (
796        "friction.auto_log",
797        &["auto log friction", "automatically log friction"],
798    ),
799    (
800        "friction.resolve",
801        &["resolve friction", "resolve this friction"],
802    ),
803    (
804        "improve.proposals",
805        &[
806            "what proposals are active",
807            "improvement proposals",
808            "list proposals",
809        ],
810    ),
811    // Claims
812    (
813        "claims",
814        &[
815            "add a claim",
816            "resolve a claim",
817            "claims status",
818            "what claims are pending",
819            "list claims",
820        ],
821    ),
822    // Transaction
823    ("transaction.begin", &["begin a transaction"]),
824    ("transaction.commit", &["commit the transaction"]),
825    ("transaction.rollback", &["rollback the transaction"]),
826    // Tools / meta
827    (
828        "tools.list",
829        &[
830            "list tools",
831            "what tools do you have",
832            "tools list",
833            "list all tools",
834        ],
835    ),
836    (
837        "nlu.shadow_report",
838        &["nlu shadow report", "show shadow mode stats"],
839    ),
840    (
841        "nlu.classify",
842        &["nlu classification test", "classify this query"],
843    ),
844    (
845        "state.snapshot",
846        &[
847            "what is the brain wave state",
848            "brain wave state",
849            "current brain wave",
850        ],
851    ),
852    (
853        "system.stats",
854        &["system stats", "show resource usage", "system stats please"],
855    ),
856    (
857        "system.health",
858        &[
859            "health check",
860            "doctor check",
861            "run a health check",
862            "system health",
863        ],
864    ),
865    (
866        "galaxy.dashboard",
867        &["consciousness dashboard", "display the dashboard"],
868    ),
869    (
870        "consciousness.depth",
871        &["consciousness depth", "depth of consciousness"],
872    ),
873    // Web / research
874    (
875        "web.fetch",
876        &[
877            "fetch this webpage",
878            "fetch the url and summarize",
879            "fetch url",
880        ],
881    ),
882    ("web.search", &["search the web for", "web search"]),
883    (
884        "web.search_and_read",
885        &["search and read", "search the web and read"],
886    ),
887    ("web.deep_fetch", &["deep fetch", "deep fetch this page"]),
888    (
889        "research.topic",
890        &[
891            "research the topic of",
892            "research topic",
893            "do a deep search on",
894        ],
895    ),
896    (
897        "research.repo",
898        &["research a github repo", "research repo", "github repo"],
899    ),
900    (
901        "research.rabbit_hole",
902        &["rabbit hole research", "rabbit hole"],
903    ),
904    // Self-play
905    (
906        "simulation.calibrate",
907        &[
908            "calibrate my predictions",
909            "record a prediction",
910            "brier scorecard",
911            "resolve a forecast",
912        ],
913    ),
914    (
915        "selfplay.run",
916        &["run selfplay", "start selfplay", "run training"],
917    ),
918    ("selfplay.status", &["selfplay status", "training status"]),
919    (
920        "selfplay.export",
921        &["export training data", "export selfplay data"],
922    ),
923    // Simulation / imagination
924    (
925        "sim.mc",
926        &[
927            "run a simulation",
928            "monte carlo simulation",
929            "simulate this",
930        ],
931    ),
932    (
933        "imagine.scenario",
934        &["imagine a scenario", "scenario planning"],
935    ),
936    (
937        "imagine.reflect",
938        &["reflect on this scenario", "counterfactual replay"],
939    ),
940    (
941        "gnosis",
942        &[
943            "what is your gana",
944            "who are you",
945            "what do I know about the wm project",
946        ],
947    ),
948];
949
950/// Merge registry tool descriptions with intent anchors for embedding.
951///
952/// Each description becomes: `"<name>: <registry description> — users say:
953/// <anchors joined>"`. Tools without anchors keep their prose description.
954///
955/// Tools whose `description()` is the Gana-level fallback (the default
956/// `Tool::description()` returns `gana().description()`, so ~45 tools across
957/// the registry embed to one of 28 shared vectors — e.g. all conformal.* and
958/// selfmodel.* tools) get a synthesized description from their dotted name:
959/// `"conformal.monitor"` → `"conformal monitor — monitor conformal
960/// prediction coverage and drift"`. Without this the router's margin
961/// calculation collapses on families with shared Gana text.
962#[must_use]
963pub fn anchored_descriptions(tools: &[Arc<dyn wm_core::Tool>]) -> Vec<(String, String)> {
964    tools
965        .iter()
966        .map(|t| {
967            let name = t.name();
968            let gana_fallback = t.gana().description() == t.description();
969            let desc = if gana_fallback {
970                synthesize_description(name)
971            } else {
972                t.description().to_string()
973            };
974            let anchors = INTENT_ANCHORS
975                .iter()
976                .find(|(n, _)| *n == name)
977                .map(|(_, a)| a);
978            let text = match anchors {
979                Some(anchors) => format!("{name}: {desc} — users say: {}", anchors.join("; ")),
980                None => format!("{name}: {desc}"),
981            };
982            (name.to_string(), text)
983        })
984        .collect()
985}
986
987/// Build a description from a dotted tool name when the tool has no explicit
988/// description (falls back to its Gana's generic text).
989///
990/// `"conformal.monitor"` → `"conformal monitor — monitor conformal prediction
991/// coverage and drift"`. The family verb (the last segment) is repeated as a
992/// verb so the embedded text carries tool-specific intent instead of the
993/// shared Gana vector.
994fn synthesize_description(name: &str) -> String {
995    let parts: Vec<&str> = name.split('.').collect();
996    if parts.len() < 2 {
997        return format!("{name} — {name} operations");
998    }
999    let family = parts[..parts.len() - 1].join(" ");
1000    let verb = parts[parts.len() - 1];
1001    let verb_hyphen = verb.replace('_', "-");
1002    format!("{family} {verb_hyphen} — {family} {verb} operations and status")
1003}
1004
1005// ── Tests ────────────────────────────────────────────────────────────
1006
1007#[cfg(test)]
1008mod tests {
1009    use super::*;
1010
1011    // --- Unit tests for helper functions ---
1012
1013    #[test]
1014    fn cosine_sim_identical_vectors() {
1015        let v = vec![1.0, 2.0, 3.0];
1016        let sim = cosine_sim(&v, &v);
1017        assert!(
1018            (sim - 1.0).abs() < 1e-5,
1019            "identical vectors should have sim=1.0, got {sim}"
1020        );
1021    }
1022
1023    #[test]
1024    fn cosine_sim_orthogonal_vectors() {
1025        let a = vec![1.0, 0.0];
1026        let b = vec![0.0, 1.0];
1027        let sim = cosine_sim(&a, &b);
1028        assert!(
1029            sim.abs() < 1e-5,
1030            "orthogonal vectors should have sim=0.0, got {sim}"
1031        );
1032    }
1033
1034    #[test]
1035    fn cosine_sim_empty_vectors() {
1036        let sim = cosine_sim(&[], &[]);
1037        assert_eq!(sim, 0.0);
1038    }
1039
1040    #[test]
1041    fn cosine_sim_different_lengths() {
1042        let a = vec![1.0, 2.0];
1043        let b = vec![1.0, 2.0, 3.0];
1044        let sim = cosine_sim(&a, &b);
1045        assert_eq!(sim, 0.0, "different-length vectors should return 0.0");
1046    }
1047
1048    #[test]
1049    fn interpolate_midpoint() {
1050        let base = vec![0.0, 0.0];
1051        let target = vec![10.0, 20.0];
1052        let result = interpolate(&base, &target, 0.5);
1053        assert!((result[0] - 5.0).abs() < 1e-5);
1054        assert!((result[1] - 10.0).abs() < 1e-5);
1055    }
1056
1057    #[test]
1058    fn interpolate_zero_alpha_returns_base() {
1059        let base = vec![1.0, 2.0, 3.0];
1060        let target = vec![10.0, 20.0, 30.0];
1061        let result = interpolate(&base, &target, 0.0);
1062        assert_eq!(result, base);
1063    }
1064
1065    #[test]
1066    fn interpolate_one_alpha_returns_target() {
1067        let base = vec![1.0, 2.0, 3.0];
1068        let target = vec![10.0, 20.0, 30.0];
1069        let result = interpolate(&base, &target, 1.0);
1070        assert_eq!(result, target);
1071    }
1072
1073    // --- OutcomeStats tests ---
1074
1075    #[test]
1076    fn outcome_stats_starts_empty() {
1077        let stats = OutcomeStats::new(384);
1078        assert_eq!(stats.success_count, 0);
1079        assert_eq!(stats.failure_count, 0);
1080        assert!(!stats.is_ready());
1081    }
1082
1083    #[test]
1084    fn outcome_stats_records_success() {
1085        let mut stats = OutcomeStats::new(4);
1086        stats.record(&[1.0, 0.0, 0.0, 0.0], true);
1087        assert_eq!(stats.success_count, 1);
1088        assert_eq!(stats.failure_count, 0);
1089    }
1090
1091    #[test]
1092    fn outcome_stats_records_failure() {
1093        let mut stats = OutcomeStats::new(4);
1094        stats.record(&[0.0, 1.0, 0.0, 0.0], false);
1095        assert_eq!(stats.success_count, 0);
1096        assert_eq!(stats.failure_count, 1);
1097    }
1098
1099    #[test]
1100    fn outcome_stats_centroid_converges() {
1101        let mut stats = OutcomeStats::new(2);
1102        // Record 3 successes at the same point
1103        for _ in 0..3 {
1104            stats.record(&[1.0, 0.0], true);
1105        }
1106        // Centroid should converge to [1.0, 0.0]
1107        assert!((stats.success_centroid[0] - 1.0).abs() < 1e-3);
1108        assert!(stats.success_centroid[1].abs() < 1e-3);
1109    }
1110
1111    #[test]
1112    fn outcome_stats_becomes_ready_after_min_observations() {
1113        let mut stats = OutcomeStats::new(2);
1114        for _ in 0..OATS_MIN_OBSERVATIONS {
1115            stats.record(&[1.0, 0.0], true);
1116        }
1117        assert!(stats.is_ready());
1118    }
1119
1120    #[test]
1121    fn outcome_stats_ignores_empty_embedding() {
1122        let mut stats = OutcomeStats::new(4);
1123        stats.record(&[], true);
1124        assert_eq!(stats.success_count, 0);
1125    }
1126
1127    // --- Tool description generation tests ---
1128
1129    #[test]
1130    fn tool_descriptions_non_empty() {
1131        let descs = tool_descriptions();
1132        assert!(
1133            !descs.is_empty(),
1134            "should have descriptions for all profiles"
1135        );
1136        assert!(
1137            descs.len() >= 60,
1138            "expected 60+ descriptions, got {}",
1139            descs.len()
1140        );
1141    }
1142
1143    #[test]
1144    fn tool_descriptions_contain_tool_name() {
1145        let descs = tool_descriptions();
1146        for (name, desc) in &descs {
1147            assert!(
1148                desc.starts_with(name),
1149                "description for '{name}' should start with the tool name, got: {desc}"
1150            );
1151        }
1152    }
1153
1154    #[test]
1155    fn tool_descriptions_contain_keywords() {
1156        let descs = tool_descriptions();
1157        let memory_create = descs.iter().find(|(n, _)| n == "memory.create");
1158        assert!(memory_create.is_some());
1159        let (_, desc) = memory_create.unwrap();
1160        assert!(
1161            desc.contains("remember"),
1162            "memory.create description should contain 'remember'"
1163        );
1164        assert!(
1165            desc.contains("store"),
1166            "memory.create description should contain 'store'"
1167        );
1168    }
1169
1170    #[test]
1171    fn tool_descriptions_are_unique() {
1172        let descs = tool_descriptions();
1173        let names: Vec<&str> = descs.iter().map(|(n, _)| n.as_str()).collect();
1174        let set: std::collections::HashSet<&str> = names.iter().copied().collect();
1175        assert_eq!(
1176            names.len(),
1177            set.len(),
1178            "duplicate tool names in descriptions"
1179        );
1180    }
1181
1182    // --- EmbeddingRouter with stub embedder ---
1183
1184    #[test]
1185    fn embedding_router_returns_none_for_stub() {
1186        let stub = Box::new(wm_memory::StubEmbedder::default());
1187        let router = EmbeddingRouter::new(stub);
1188        assert!(
1189            router.is_none(),
1190            "embedding router should return None for stub embedder"
1191        );
1192    }
1193
1194    #[test]
1195    fn embedding_router_with_descriptions_covers_registry_tools() {
1196        let embedder = Box::new(KeywordEmbedder::new(vec![
1197            "memory", "karma", "session", "list",
1198        ]));
1199        let descriptions = vec![
1200            (
1201                "memory.create".to_string(),
1202                "remember and store information in persistent memory".to_string(),
1203            ),
1204            (
1205                "karma.clear".to_string(),
1206                "wipe and reset the karma ledger entries".to_string(),
1207            ),
1208            (
1209                "session.list".to_string(),
1210                "list all recorded sessions".to_string(),
1211            ),
1212        ];
1213        let router =
1214            EmbeddingRouter::with_descriptions(embedder, descriptions).expect("should init");
1215        assert_eq!(router.tool_count(), 3);
1216        let (tool, _) = router.route("show me the sessions");
1217        assert_eq!(
1218            tool, "session.list",
1219            "registry-description routing should find session.list"
1220        );
1221    }
1222
1223    #[test]
1224    fn route_with_margin_returns_positive_margin() {
1225        let keywords: Vec<&str> = TOOL_PROFILES
1226            .iter()
1227            .flat_map(|p| p.keywords.iter().map(|(t, _)| *t))
1228            .collect::<std::collections::HashSet<_>>()
1229            .into_iter()
1230            .collect();
1231        let embedder = Box::new(KeywordEmbedder::new(keywords));
1232        let router = EmbeddingRouter::new(embedder).expect("should init");
1233
1234        let (tool, conf, margin) = router
1235            .route_with_margin("remember that the sky is blue")
1236            .expect("clear match should return Some");
1237        assert_eq!(tool, "memory.create");
1238        assert!(conf > 0.0);
1239        assert!(margin >= 0.0, "margin should be non-negative");
1240    }
1241
1242    // --- Mock embedder for testing ---
1243
1244    /// A test embedder that generates simple keyword-based embeddings.
1245    /// Each dimension corresponds to a keyword — if the text contains the
1246    /// keyword, that dimension is 1.0, otherwise 0.0. This provides basic
1247    /// semantic similarity for testing without a real embedder.
1248    struct KeywordEmbedder {
1249        keywords: Vec<String>,
1250        dim: usize,
1251    }
1252
1253    impl KeywordEmbedder {
1254        fn new(keywords: Vec<&str>) -> Self {
1255            let dim = keywords.len();
1256            Self {
1257                keywords: keywords.into_iter().map(String::from).collect(),
1258                dim,
1259            }
1260        }
1261
1262        fn embed_text(&self, text: &str) -> Vec<f32> {
1263            let lower = text.to_lowercase();
1264            self.keywords
1265                .iter()
1266                .map(|kw| {
1267                    if lower.contains(&kw.to_lowercase()) {
1268                        1.0
1269                    } else {
1270                        0.0
1271                    }
1272                })
1273                .collect()
1274        }
1275    }
1276
1277    impl Embedder for KeywordEmbedder {
1278        fn embed_batch(&self, texts: &[&str]) -> wm_core::Result<Vec<Vec<f32>>> {
1279            Ok(texts.iter().map(|t| self.embed_text(t)).collect())
1280        }
1281
1282        fn dimension(&self) -> usize {
1283            self.dim
1284        }
1285
1286        fn is_available(&self) -> bool {
1287            true
1288        }
1289
1290        fn backend_name(&self) -> &'static str {
1291            "keyword-test"
1292        }
1293    }
1294
1295    #[test]
1296    fn embedding_router_works_with_keyword_embedder() {
1297        let keywords: Vec<&str> = TOOL_PROFILES
1298            .iter()
1299            .flat_map(|p| p.keywords.iter().map(|(t, _)| *t))
1300            .collect::<std::collections::HashSet<_>>()
1301            .into_iter()
1302            .collect();
1303        let embedder = Box::new(KeywordEmbedder::new(keywords));
1304        let router = EmbeddingRouter::new(embedder).expect("should init with keyword embedder");
1305
1306        assert!(router.tool_count() >= 60);
1307        assert!(router.dimension() > 0);
1308        assert_eq!(router.backend_name(), "keyword-test");
1309    }
1310
1311    #[test]
1312    fn embedding_router_routes_remember_to_memory_create() {
1313        let keywords: Vec<&str> = TOOL_PROFILES
1314            .iter()
1315            .flat_map(|p| p.keywords.iter().map(|(t, _)| *t))
1316            .collect::<std::collections::HashSet<_>>()
1317            .into_iter()
1318            .collect();
1319        let embedder = Box::new(KeywordEmbedder::new(keywords));
1320        let router = EmbeddingRouter::new(embedder).expect("should init");
1321
1322        let (tool, conf) = router.route("remember that the sky is blue");
1323        assert_eq!(tool, "memory.create");
1324        assert!(
1325            conf > 0.0,
1326            "confidence should be > 0 for clear match, got {conf}"
1327        );
1328    }
1329
1330    #[test]
1331    fn embedding_router_routes_search_to_memory_search() {
1332        let keywords: Vec<&str> = TOOL_PROFILES
1333            .iter()
1334            .flat_map(|p| p.keywords.iter().map(|(t, _)| *t))
1335            .collect::<std::collections::HashSet<_>>()
1336            .into_iter()
1337            .collect();
1338        let embedder = Box::new(KeywordEmbedder::new(keywords));
1339        let router = EmbeddingRouter::new(embedder).expect("should init");
1340
1341        let (tool, conf) = router.route("search for rust");
1342        assert_eq!(tool, "memory.search");
1343        assert!(conf > 0.0);
1344    }
1345
1346    #[test]
1347    fn embedding_router_empty_returns_gnosis() {
1348        let keywords: Vec<&str> = TOOL_PROFILES
1349            .iter()
1350            .flat_map(|p| p.keywords.iter().map(|(t, _)| *t))
1351            .collect::<std::collections::HashSet<_>>()
1352            .into_iter()
1353            .collect();
1354        let embedder = Box::new(KeywordEmbedder::new(keywords));
1355        let router = EmbeddingRouter::new(embedder).expect("should init");
1356
1357        let (tool, conf) = router.route("");
1358        assert_eq!(tool, "gnosis");
1359        assert_eq!(conf, 0.0);
1360    }
1361
1362    #[test]
1363    fn embedding_router_whitespace_returns_gnosis() {
1364        let keywords: Vec<&str> = TOOL_PROFILES
1365            .iter()
1366            .flat_map(|p| p.keywords.iter().map(|(t, _)| *t))
1367            .collect::<std::collections::HashSet<_>>()
1368            .into_iter()
1369            .collect();
1370        let embedder = Box::new(KeywordEmbedder::new(keywords));
1371        let router = EmbeddingRouter::new(embedder).expect("should init");
1372
1373        let (tool, conf) = router.route("   ");
1374        assert_eq!(tool, "gnosis");
1375        assert_eq!(conf, 0.0);
1376    }
1377
1378    #[test]
1379    fn embedding_router_unknown_returns_gnosis() {
1380        let keywords: Vec<&str> = TOOL_PROFILES
1381            .iter()
1382            .flat_map(|p| p.keywords.iter().map(|(t, _)| *t))
1383            .collect::<std::collections::HashSet<_>>()
1384            .into_iter()
1385            .collect();
1386        let embedder = Box::new(KeywordEmbedder::new(keywords));
1387        let router = EmbeddingRouter::new(embedder).expect("should init");
1388
1389        let (tool, _conf) = router.route("xyzzy frobnicate");
1390        assert_eq!(tool, "gnosis");
1391    }
1392
1393    #[test]
1394    fn embedding_router_record_outcome_updates_stats() {
1395        let keywords: Vec<&str> = TOOL_PROFILES
1396            .iter()
1397            .flat_map(|p| p.keywords.iter().map(|(t, _)| *t))
1398            .collect::<std::collections::HashSet<_>>()
1399            .into_iter()
1400            .collect();
1401        let embedder = Box::new(KeywordEmbedder::new(keywords));
1402        let router = EmbeddingRouter::new(embedder).expect("should init");
1403
1404        // Record some outcomes
1405        router.record_outcome("memory.create", "remember that rust is fast", true);
1406        router.record_outcome("memory.create", "store this fact", true);
1407        router.record_outcome("memory.search", "search for rust", false);
1408
1409        let counts = router.outcome_counts();
1410        let memory_create = counts.iter().find(|(n, _, _)| n == "memory.create");
1411        assert!(memory_create.is_some());
1412        let (_, success, failure) = memory_create.unwrap();
1413        assert_eq!(*success, 2);
1414        assert_eq!(*failure, 0);
1415    }
1416
1417    #[test]
1418    fn embedding_router_record_outcome_ignores_empty_query() {
1419        let keywords: Vec<&str> = TOOL_PROFILES
1420            .iter()
1421            .flat_map(|p| p.keywords.iter().map(|(t, _)| *t))
1422            .collect::<std::collections::HashSet<_>>()
1423            .into_iter()
1424            .collect();
1425        let embedder = Box::new(KeywordEmbedder::new(keywords));
1426        let router = EmbeddingRouter::new(embedder).expect("should init");
1427
1428        router.record_outcome("memory.create", "", true);
1429        let counts = router.outcome_counts();
1430        assert!(
1431            counts.is_empty(),
1432            "empty query should not create outcome stats"
1433        );
1434    }
1435
1436    #[test]
1437    fn embedding_router_oats_refine_improves_routing() {
1438        let keywords: Vec<&str> = TOOL_PROFILES
1439            .iter()
1440            .flat_map(|p| p.keywords.iter().map(|(t, _)| *t))
1441            .collect::<std::collections::HashSet<_>>()
1442            .into_iter()
1443            .collect();
1444        let embedder = Box::new(KeywordEmbedder::new(keywords));
1445        let router = EmbeddingRouter::new(embedder).expect("should init");
1446
1447        // Record many successes for memory.create with "save" queries
1448        for _ in 0..15 {
1449            router.record_outcome("memory.create", "save this important fact", true);
1450        }
1451
1452        // Now "save this important fact" should route to memory.create with high confidence
1453        let (tool, conf) = router.route("save this important fact");
1454        assert_eq!(tool, "memory.create");
1455        assert!(
1456            conf > 0.0,
1457            "OATS-refined routing should still match, got conf={conf}"
1458        );
1459    }
1460
1461    // --- A/B comparison: embedding router vs TF-IDF on key test cases ---
1462
1463    #[test]
1464    fn ab_comparison_remember() {
1465        let keywords: Vec<&str> = TOOL_PROFILES
1466            .iter()
1467            .flat_map(|p| p.keywords.iter().map(|(t, _)| *t))
1468            .collect::<std::collections::HashSet<_>>()
1469            .into_iter()
1470            .collect();
1471        let embedder = Box::new(KeywordEmbedder::new(keywords));
1472        let router = EmbeddingRouter::new(embedder).expect("should init");
1473
1474        let query = "remember that the sky is blue";
1475        let (emb_tool, emb_conf) = router.route(query);
1476        let (tfidf_tool, tfidf_conf) = crate::nlu::classify(query);
1477
1478        assert_eq!(
1479            emb_tool, tfidf_tool,
1480            "embedding and TF-IDF should agree on '{query}'"
1481        );
1482        assert!(emb_conf > 0.0 && tfidf_conf > 0.0);
1483    }
1484
1485    #[test]
1486    fn ab_comparison_search() {
1487        let keywords: Vec<&str> = TOOL_PROFILES
1488            .iter()
1489            .flat_map(|p| p.keywords.iter().map(|(t, _)| *t))
1490            .collect::<std::collections::HashSet<_>>()
1491            .into_iter()
1492            .collect();
1493        let embedder = Box::new(KeywordEmbedder::new(keywords));
1494        let router = EmbeddingRouter::new(embedder).expect("should init");
1495
1496        let query = "search for rust";
1497        let (emb_tool, emb_conf) = router.route(query);
1498        let (tfidf_tool, _) = crate::nlu::classify(query);
1499
1500        assert_eq!(
1501            emb_tool, tfidf_tool,
1502            "embedding and TF-IDF should agree on '{query}'"
1503        );
1504        assert!(emb_conf > 0.0);
1505    }
1506
1507    #[test]
1508    fn ab_comparison_delete() {
1509        let keywords: Vec<&str> = TOOL_PROFILES
1510            .iter()
1511            .flat_map(|p| p.keywords.iter().map(|(t, _)| *t))
1512            .collect::<std::collections::HashSet<_>>()
1513            .into_iter()
1514            .collect();
1515        let embedder = Box::new(KeywordEmbedder::new(keywords));
1516        let router = EmbeddingRouter::new(embedder).expect("should init");
1517
1518        let query = "delete memory abc-123";
1519        let (emb_tool, _) = router.route(query);
1520        let (tfidf_tool, _) = crate::nlu::classify(query);
1521
1522        assert_eq!(
1523            emb_tool, tfidf_tool,
1524            "embedding and TF-IDF should agree on '{query}'"
1525        );
1526    }
1527
1528    #[test]
1529    fn ab_comparison_karma() {
1530        let keywords: Vec<&str> = TOOL_PROFILES
1531            .iter()
1532            .flat_map(|p| p.keywords.iter().map(|(t, _)| *t))
1533            .collect::<std::collections::HashSet<_>>()
1534            .into_iter()
1535            .collect();
1536        let embedder = Box::new(KeywordEmbedder::new(keywords));
1537        let router = EmbeddingRouter::new(embedder).expect("should init");
1538
1539        let query = "show me the karma report";
1540        let (emb_tool, _) = router.route(query);
1541        let (tfidf_tool, _) = crate::nlu::classify(query);
1542
1543        assert_eq!(
1544            emb_tool, tfidf_tool,
1545            "embedding and TF-IDF should agree on '{query}'"
1546        );
1547    }
1548
1549    // ── ShadowModeStats tests ─────────────────────────────────────────
1550
1551    #[test]
1552    fn shadow_stats_record_agreement() {
1553        let mut stats = ShadowModeStats::default();
1554        stats.record("test query", "memory.create", 0.9, "memory.create", 0.8);
1555        assert_eq!(stats.total_queries, 1);
1556        assert_eq!(stats.total_disagreements, 0);
1557        assert!(stats.samples.is_empty());
1558    }
1559
1560    #[test]
1561    fn shadow_stats_record_disagreement() {
1562        let mut stats = ShadowModeStats::default();
1563        stats.record("test query", "memory.create", 0.9, "memory.list", 0.7);
1564        assert_eq!(stats.total_queries, 1);
1565        assert_eq!(stats.total_disagreements, 1);
1566        assert_eq!(stats.samples.len(), 1);
1567        assert_eq!(stats.samples[0].embedding_tool, "memory.create");
1568        assert_eq!(stats.samples[0].tfidf_tool, "memory.list");
1569    }
1570
1571    #[test]
1572    fn shadow_stats_disagreement_rate() {
1573        let mut stats = ShadowModeStats::default();
1574        for _ in 0..8 {
1575            stats.record("agree", "memory.create", 0.9, "memory.create", 0.8);
1576        }
1577        for _ in 0..2 {
1578            stats.record("disagree", "memory.create", 0.9, "memory.list", 0.7);
1579        }
1580        assert_eq!(stats.total_queries, 10);
1581        assert_eq!(stats.total_disagreements, 2);
1582        assert!((stats.disagreement_rate() - 0.2).abs() < 0.001);
1583    }
1584
1585    #[test]
1586    fn shadow_stats_promotion_ready_threshold() {
1587        let mut stats = ShadowModeStats::default();
1588        // Not enough queries
1589        for _ in 0..99 {
1590            stats.record("agree", "memory.create", 0.9, "memory.create", 0.8);
1591        }
1592        assert!(!stats.promotion_ready());
1593
1594        // Enough queries, low disagreement
1595        stats.record("agree", "memory.create", 0.9, "memory.create", 0.8);
1596        assert!(stats.promotion_ready());
1597
1598        // Too many disagreements (25 out of 125 = 0.20, not < 0.20)
1599        for _ in 0..25 {
1600            stats.record("disagree", "memory.create", 0.9, "memory.list", 0.7);
1601        }
1602        assert!(!stats.promotion_ready());
1603    }
1604
1605    #[test]
1606    fn shadow_stats_report_json() {
1607        let mut stats = ShadowModeStats::default();
1608        stats.record("test", "memory.create", 0.9, "memory.list", 0.7);
1609        let report = stats.report();
1610        assert_eq!(report["total_queries"], 1);
1611        assert_eq!(report["total_disagreements"], 1);
1612        assert!(report["promotion_ready"].is_boolean());
1613        assert!(report["recent_samples"].is_array());
1614    }
1615
1616    #[test]
1617    fn shadow_stats_samples_capped() {
1618        let mut stats = ShadowModeStats::default();
1619        for i in 0..100 {
1620            stats.record(
1621                &format!("query {i}"),
1622                "memory.create",
1623                0.9,
1624                "memory.list",
1625                0.7,
1626            );
1627        }
1628        assert_eq!(stats.samples.len(), 50); // MAX_SAMPLES
1629    }
1630
1631    #[test]
1632    fn shadow_stats_serialization_roundtrip() {
1633        let mut stats = ShadowModeStats::default();
1634        stats.record("test", "memory.create", 0.9, "memory.list", 0.7);
1635        stats.record("another", "gnosis", 0.1, "gnosis", 0.1);
1636        let json = serde_json::to_string(&stats).unwrap();
1637        let deserialized: ShadowModeStats = serde_json::from_str(&json).unwrap();
1638        assert_eq!(deserialized.total_queries, 2);
1639        assert_eq!(deserialized.total_disagreements, 1);
1640        assert_eq!(deserialized.samples.len(), 1);
1641    }
1642
1643    #[test]
1644    fn oats_persistence_roundtrip() {
1645        let keywords: Vec<&str> = TOOL_PROFILES
1646            .iter()
1647            .flat_map(|p| p.keywords.iter().map(|(t, _)| *t))
1648            .collect::<std::collections::HashSet<_>>()
1649            .into_iter()
1650            .collect();
1651        let embedder = Box::new(KeywordEmbedder::new(keywords));
1652        let router = EmbeddingRouter::new(embedder).expect("should init");
1653
1654        // Record some outcomes
1655        router.record_outcome("memory.create", "create a memory", true);
1656        router.record_outcome("memory.create", "store this", true);
1657        router.record_outcome("memory.list", "list memories", true);
1658
1659        // Save
1660        let saved = router.save_oats().expect("should serialize");
1661
1662        // Load into a new router
1663        let keywords2: Vec<&str> = TOOL_PROFILES
1664            .iter()
1665            .flat_map(|p| p.keywords.iter().map(|(t, _)| *t))
1666            .collect::<std::collections::HashSet<_>>()
1667            .into_iter()
1668            .collect();
1669        let embedder2 = Box::new(KeywordEmbedder::new(keywords2));
1670        let router2 = EmbeddingRouter::new(embedder2).expect("should init");
1671        router2.load_oats(&saved);
1672
1673        let counts1 = router.outcome_counts();
1674        let counts2 = router2.outcome_counts();
1675        assert_eq!(counts1.len(), counts2.len());
1676        for (name, success, failure) in &counts1 {
1677            let match_found = counts2
1678                .iter()
1679                .any(|(n, s, f)| n == name && s == success && f == failure);
1680            assert!(
1681                match_found,
1682                "OATS data should match after roundtrip for {name}"
1683            );
1684        }
1685    }
1686}