1use ahash::AHashMap;
28use std::sync::{Arc, RwLock};
29use wm_memory::Embedder;
30
31use crate::nlu::{PREFIX_ROUTES, TOOL_PROFILES, ToolProfile};
32
33const OATS_ALPHA: f32 = 0.15;
35
36const OATS_MIN_OBSERVATIONS: usize = 10;
38
39const MIN_THRESHOLD: f64 = 0.10;
41
42pub const MIN_MARGIN: f64 = 0.02;
50
51#[derive(Debug, Clone)]
53pub struct OutcomeStats {
54 success_centroid: Vec<f32>,
56 #[allow(dead_code)]
58 failure_centroid: Vec<f32>,
59 success_count: usize,
61 failure_count: usize,
63}
64
65impl OutcomeStats {
66 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 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 const fn is_ready(&self) -> bool {
99 self.success_count >= OATS_MIN_OBSERVATIONS
100 }
101}
102
103fn 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
115pub struct EmbeddingRouter {
122 tool_embeddings: AHashMap<String, Vec<f32>>,
124 embedder: Box<dyn Embedder>,
126 outcome_stats: RwLock<AHashMap<String, OutcomeStats>>,
128 dim: usize,
130 apply_prefix_bonus: bool,
138}
139
140impl EmbeddingRouter {
141 #[must_use]
149 pub fn new(embedder: Box<dyn Embedder>) -> Option<Self> {
150 Self::new_with_descriptions(embedder, tool_descriptions(), true)
151 }
152
153 #[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 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 #[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 #[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 #[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 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 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 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 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 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 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 #[must_use]
395 pub fn tool_count(&self) -> usize {
396 self.tool_embeddings.len()
397 }
398
399 #[must_use]
401 pub const fn dimension(&self) -> usize {
402 self.dim
403 }
404
405 #[must_use]
407 pub fn backend_name(&self) -> &str {
408 self.embedder.backend_name()
409 }
410
411 #[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 #[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 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
467const MAX_SAMPLES: usize = 50;
471
472#[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#[derive(Debug, Default, serde::Serialize, serde::Deserialize)]
487pub struct ShadowModeStats {
488 pub total_queries: u64,
490 pub total_disagreements: u64,
492 pub disagreement_pairs: std::collections::HashMap<String, u64>,
494 pub samples: Vec<DisagreementSample>,
496}
497
498impl ShadowModeStats {
499 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 #[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 #[must_use]
539 pub fn promotion_ready(&self) -> bool {
540 self.total_queries >= 100 && self.disagreement_rate() < 0.20
541 }
542
543 #[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
574fn 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
596fn 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#[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
618fn 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
624static INTENT_ANCHORS: &[(&str, &[&str])] = &[
632 (
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 (
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.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 (
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 (
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 (
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.begin", &["begin a transaction"]),
824 ("transaction.commit", &["commit the transaction"]),
825 ("transaction.rollback", &["rollback the transaction"]),
826 (
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 (
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 (
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 (
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#[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
987fn 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#[cfg(test)]
1008mod tests {
1009 use super::*;
1010
1011 #[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 #[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 for _ in 0..3 {
1104 stats.record(&[1.0, 0.0], true);
1105 }
1106 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 #[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 #[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 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 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 for _ in 0..15 {
1449 router.record_outcome("memory.create", "save this important fact", true);
1450 }
1451
1452 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 #[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 #[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 for _ in 0..99 {
1590 stats.record("agree", "memory.create", 0.9, "memory.create", 0.8);
1591 }
1592 assert!(!stats.promotion_ready());
1593
1594 stats.record("agree", "memory.create", 0.9, "memory.create", 0.8);
1596 assert!(stats.promotion_ready());
1597
1598 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); }
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 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 let saved = router.save_oats().expect("should serialize");
1661
1662 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}