use std::collections::{HashMap, HashSet};
use std::sync::Arc;
use std::time::Instant;
use anyhow::Result;
use chrono::{DateTime, Utc};
use parking_lot::RwLock;
use regex::Regex;
use serde::{Deserialize, Serialize};
use uuid::Uuid;
use crate::embeddings::NeuralNer;
use crate::graph_memory::GraphMemory;
use crate::memory::feedback::FeedbackStore;
use crate::memory::{Memory, MemorySystem};
fn contains_word(text: &str, word: &str) -> bool {
if word.is_empty() {
return false;
}
let escaped = regex::escape(word);
let pattern = format!(r"(?i)\b{}\b", escaped);
match Regex::new(&pattern) {
Ok(re) => re.is_match(text),
Err(_) => text.contains(word), }
}
const DEFAULT_SEMANTIC_WEIGHT: f32 = 0.18;
const DEFAULT_ENTITY_WEIGHT: f32 = 0.17;
const DEFAULT_TAG_WEIGHT: f32 = 0.05;
const DEFAULT_IMPORTANCE_WEIGHT: f32 = 0.05;
const DEFAULT_MOMENTUM_WEIGHT: f32 = 0.28;
const DEFAULT_ACCESS_COUNT_WEIGHT: f32 = 0.14;
const DEFAULT_GRAPH_STRENGTH_WEIGHT: f32 = 0.13;
const WEIGHT_LEARNING_RATE: f32 = 0.05;
const MIN_WEIGHT: f32 = 0.05;
const SIGMOID_STEEPNESS: f32 = 10.0;
const SIGMOID_MIDPOINT: f32 = 0.5;
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct RelevanceConfig {
#[serde(default = "default_semantic_threshold")]
pub semantic_threshold: f32,
#[serde(default = "default_entity_threshold")]
pub entity_threshold: f32,
#[serde(default = "default_max_results")]
pub max_results: usize,
#[serde(default)]
pub memory_types: Vec<String>,
#[serde(default = "default_true")]
pub enable_entity_matching: bool,
#[serde(default = "default_true")]
pub enable_semantic_matching: bool,
#[serde(default = "default_min_importance")]
pub min_importance: f32,
#[serde(default = "default_recency_hours")]
pub recency_boost_hours: u64,
#[serde(default = "default_recency_multiplier")]
pub recency_boost_multiplier: f32,
#[serde(default = "default_graph_boost_multiplier")]
pub graph_boost_multiplier: f32,
}
fn default_graph_boost_multiplier() -> f32 {
1.15 }
fn default_semantic_threshold() -> f32 {
0.45 }
fn default_entity_threshold() -> f32 {
0.5
}
fn default_max_results() -> usize {
5
}
fn default_true() -> bool {
true
}
fn default_min_importance() -> f32 {
0.3
}
fn default_recency_hours() -> u64 {
24
}
fn default_recency_multiplier() -> f32 {
1.2
}
impl Default for RelevanceConfig {
fn default() -> Self {
Self {
semantic_threshold: default_semantic_threshold(),
entity_threshold: default_entity_threshold(),
max_results: default_max_results(),
memory_types: Vec::new(),
enable_entity_matching: true,
enable_semantic_matching: true,
min_importance: default_min_importance(),
recency_boost_hours: default_recency_hours(),
recency_boost_multiplier: default_recency_multiplier(),
graph_boost_multiplier: default_graph_boost_multiplier(),
}
}
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct SurfacedMemory {
pub id: String,
pub content: String,
pub memory_type: String,
pub importance: f32,
pub relevance_score: f32,
pub relevance_reason: RelevanceReason,
pub matched_entities: Vec<String>,
pub semantic_similarity: Option<f32>,
pub created_at: DateTime<Utc>,
pub tags: Vec<String>,
}
#[derive(Debug, Clone, Serialize, Deserialize, PartialEq)]
pub enum RelevanceReason {
EntityMatch,
SemanticSimilarity,
Combined,
RecentImportant,
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct RelevanceRequest {
pub user_id: String,
pub context: String,
#[serde(default)]
pub entities: Vec<String>,
#[serde(default)]
pub config: RelevanceConfig,
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct RelevanceResponse {
pub memories: Vec<SurfacedMemory>,
pub detected_entities: Vec<DetectedEntity>,
pub latency_ms: f64,
pub latency_target_met: bool,
#[serde(skip_serializing_if = "Option::is_none")]
pub debug: Option<RelevanceDebug>,
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct DetectedEntity {
pub name: String,
pub entity_type: String,
pub confidence: f32,
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct RelevanceDebug {
pub ner_ms: f64,
pub entity_match_ms: f64,
pub semantic_search_ms: f64,
pub ranking_ms: f64,
pub memories_scanned: usize,
pub entity_matches: usize,
pub semantic_matches: usize,
}
#[derive(Debug, Clone, Default)]
struct EntityIndexEntry {
memory_ids: HashSet<Uuid>,
#[allow(dead_code)]
last_updated: Option<DateTime<Utc>>,
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct LearnedWeights {
pub semantic: f32,
pub entity: f32,
pub tag: f32,
pub importance: f32,
#[serde(default = "default_momentum_weight")]
pub momentum: f32,
#[serde(default = "default_access_count_weight")]
pub access_count: f32,
#[serde(default = "default_graph_strength_weight")]
pub graph_strength: f32,
pub update_count: u32,
pub last_updated: Option<DateTime<Utc>>,
}
fn default_momentum_weight() -> f32 {
DEFAULT_MOMENTUM_WEIGHT
}
fn default_access_count_weight() -> f32 {
DEFAULT_ACCESS_COUNT_WEIGHT
}
fn default_graph_strength_weight() -> f32 {
DEFAULT_GRAPH_STRENGTH_WEIGHT
}
impl Default for LearnedWeights {
fn default() -> Self {
Self {
semantic: DEFAULT_SEMANTIC_WEIGHT,
entity: DEFAULT_ENTITY_WEIGHT,
tag: DEFAULT_TAG_WEIGHT,
importance: DEFAULT_IMPORTANCE_WEIGHT,
momentum: DEFAULT_MOMENTUM_WEIGHT,
access_count: DEFAULT_ACCESS_COUNT_WEIGHT,
graph_strength: DEFAULT_GRAPH_STRENGTH_WEIGHT,
update_count: 0,
last_updated: None,
}
}
}
impl LearnedWeights {
pub fn normalize(&mut self) {
let sum = self.semantic
+ self.entity
+ self.tag
+ self.importance
+ self.momentum
+ self.access_count
+ self.graph_strength;
if sum > 0.0 {
self.semantic /= sum;
self.entity /= sum;
self.tag /= sum;
self.importance /= sum;
self.momentum /= sum;
self.access_count /= sum;
self.graph_strength /= sum;
}
}
pub fn apply_feedback(
&mut self,
semantic_contributed: bool,
entity_contributed: bool,
tag_contributed: bool,
helpful: bool,
) {
let direction = if helpful { 1.0 } else { -1.0 };
let delta = WEIGHT_LEARNING_RATE * direction;
if semantic_contributed {
self.semantic = (self.semantic + delta).max(MIN_WEIGHT);
}
if entity_contributed {
self.entity = (self.entity + delta).max(MIN_WEIGHT);
}
if tag_contributed {
self.tag = (self.tag + delta).max(MIN_WEIGHT);
}
if helpful && !semantic_contributed && !entity_contributed && !tag_contributed {
self.importance = (self.importance + delta).max(MIN_WEIGHT);
}
let aux_delta = WEIGHT_LEARNING_RATE * direction * 0.5;
self.momentum = (self.momentum + aux_delta).max(MIN_WEIGHT);
self.access_count = (self.access_count + aux_delta).max(MIN_WEIGHT);
self.graph_strength = (self.graph_strength + aux_delta).max(MIN_WEIGHT);
self.normalize();
self.update_count += 1;
self.last_updated = Some(Utc::now());
}
pub fn fuse_scores(
&self,
semantic_score: f32,
entity_score: f32,
tag_score: f32,
importance_score: f32,
) -> f32 {
self.fuse_scores_full(
semantic_score,
entity_score,
tag_score,
importance_score,
0.0,
0,
0.5,
)
}
pub fn fuse_scores_with_momentum(
&self,
semantic_score: f32,
entity_score: f32,
tag_score: f32,
importance_score: f32,
momentum_ema: f32,
) -> f32 {
self.fuse_scores_full(
semantic_score,
entity_score,
tag_score,
importance_score,
momentum_ema,
0,
0.5,
)
}
pub fn fuse_scores_full(
&self,
semantic_score: f32,
entity_score: f32,
tag_score: f32,
importance_score: f32,
momentum_ema: f32,
access_count: u32,
graph_strength: f32,
) -> f32 {
let calibrated_semantic = calibrate_score(semantic_score);
let calibrated_entity = calibrate_score(entity_score);
let calibrated_tag = calibrate_score(tag_score);
let calibrated_importance = calibrate_score(importance_score);
let normalized_momentum = (momentum_ema + 1.0) / 2.0;
let amplified_momentum = if normalized_momentum > 0.65 {
(normalized_momentum * 1.5).min(1.0)
} else if normalized_momentum < 0.40 {
(normalized_momentum * 0.3).max(0.0)
} else {
normalized_momentum
};
let calibrated_momentum = calibrate_score(amplified_momentum);
let access_score = if access_count == 0 {
0.0
} else {
let log_access = (access_count as f32 + 1.0).log2();
(log_access / 4.0).min(1.0)
};
let calibrated_access = calibrate_score(access_score);
let calibrated_graph = calibrate_score(graph_strength);
let result = self.semantic * calibrated_semantic
+ self.entity * calibrated_entity
+ self.tag * calibrated_tag
+ self.importance * calibrated_importance
+ self.momentum * calibrated_momentum
+ self.access_count * calibrated_access
+ self.graph_strength * calibrated_graph;
if result.is_finite() {
result
} else {
0.0
}
}
}
fn calibrate_score(score: f32) -> f32 {
if !score.is_finite() {
return 0.0;
}
1.0 / (1.0 + (-SIGMOID_STEEPNESS * (score - SIGMOID_MIDPOINT)).exp())
}
pub struct RelevanceEngine {
ner: Arc<NeuralNer>,
entity_index: Arc<RwLock<HashMap<String, EntityIndexEntry>>>,
entity_index_timestamp: Arc<RwLock<Option<DateTime<Utc>>>>,
learned_weights: Arc<RwLock<LearnedWeights>>,
active_ab_test: Arc<RwLock<Option<String>>>,
}
impl RelevanceEngine {
pub fn new(ner: Arc<NeuralNer>) -> Self {
Self {
ner,
entity_index: Arc::new(RwLock::new(HashMap::new())),
entity_index_timestamp: Arc::new(RwLock::new(None)),
learned_weights: Arc::new(RwLock::new(LearnedWeights::default())),
active_ab_test: Arc::new(RwLock::new(None)),
}
}
pub fn set_active_ab_test(&self, test_id: Option<String>) {
*self.active_ab_test.write() = test_id;
}
pub fn get_active_ab_test(&self) -> Option<String> {
self.active_ab_test.read().clone()
}
pub fn get_weights(&self) -> LearnedWeights {
self.learned_weights.read().clone()
}
pub fn set_weights(&self, weights: LearnedWeights) {
*self.learned_weights.write() = weights;
}
pub fn apply_feedback(
&self,
semantic_contributed: bool,
entity_contributed: bool,
tag_contributed: bool,
helpful: bool,
) {
self.learned_weights.write().apply_feedback(
semantic_contributed,
entity_contributed,
tag_contributed,
helpful,
);
}
fn calculate_tag_score(&self, context: &str, tags: &[String]) -> f32 {
if tags.is_empty() {
return 0.0;
}
let context_lower = context.to_lowercase();
let mut matches = 0;
for tag in tags {
let tag_lower = tag.to_lowercase();
if context_lower.contains(&tag_lower) {
matches += 1;
} else {
for word in context_lower.split_whitespace() {
if word.starts_with(&tag_lower) || tag_lower.starts_with(word) {
matches += 1;
break;
}
}
}
}
matches as f32 / tags.len() as f32
}
pub fn surface_relevant(
&self,
context: &str,
memory_system: &MemorySystem,
graph_memory: Option<&GraphMemory>,
config: &RelevanceConfig,
feedback_store: Option<&RwLock<FeedbackStore>>,
) -> Result<RelevanceResponse> {
self.surface_relevant_inner(
context,
memory_system,
graph_memory,
config,
feedback_store,
None,
)
}
fn surface_relevant_inner(
&self,
context: &str,
memory_system: &MemorySystem,
graph_memory: Option<&GraphMemory>,
config: &RelevanceConfig,
feedback_store: Option<&RwLock<FeedbackStore>>,
weights_override: Option<LearnedWeights>,
) -> Result<RelevanceResponse> {
let start = Instant::now();
let mut debug = RelevanceDebug {
ner_ms: 0.0,
entity_match_ms: 0.0,
semantic_search_ms: 0.0,
ranking_ms: 0.0,
memories_scanned: 0,
entity_matches: 0,
semantic_matches: 0,
};
let ner_start = Instant::now();
let detected_entities = if config.enable_entity_matching {
self.extract_entities(context)
} else {
Vec::new()
};
debug.ner_ms = ner_start.elapsed().as_secs_f64() * 1000.0;
let mut candidate_memories: HashMap<Uuid, (Memory, f32, f32, Vec<String>)> = HashMap::new();
if config.enable_entity_matching && !detected_entities.is_empty() {
let entity_start = Instant::now();
let entity_matches =
self.match_by_entities(&detected_entities, memory_system, graph_memory, config)?;
debug.entity_match_ms = entity_start.elapsed().as_secs_f64() * 1000.0;
debug.entity_matches = entity_matches.len();
for (memory, score, matched) in entity_matches {
let id = memory.id.0;
candidate_memories.insert(id, (memory, 0.0, score, matched));
}
}
if config.enable_semantic_matching {
let semantic_start = Instant::now();
let semantic_matches = self.match_by_semantic(context, memory_system, config)?;
debug.semantic_search_ms = semantic_start.elapsed().as_secs_f64() * 1000.0;
debug.semantic_matches = semantic_matches.len();
for (memory, score) in semantic_matches {
let id = memory.id.0;
if let Some((_, semantic_score, _entity_score, _matched)) =
candidate_memories.get_mut(&id)
{
*semantic_score = score;
} else {
candidate_memories.insert(id, (memory, score, 0.0, Vec::new()));
}
}
}
debug.memories_scanned = candidate_memories.len();
let ranking_start = Instant::now();
let weights = weights_override.unwrap_or_else(|| self.learned_weights.read().clone());
let mut results: Vec<SurfacedMemory> = candidate_memories
.into_iter()
.filter_map(
|(_, (memory, semantic_score, entity_score, matched_entities))| {
let importance = memory.importance();
if importance < config.min_importance {
return None;
}
if !config.memory_types.is_empty() {
let mem_type = format!("{:?}", memory.experience.experience_type);
if !config
.memory_types
.iter()
.any(|t| t.eq_ignore_ascii_case(&mem_type))
{
return None;
}
}
let tag_score = self.calculate_tag_score(context, &memory.experience.tags);
let access_count = memory.access_count();
let graph_strength = graph_memory
.and_then(|g| g.get_memory_hebbian_strength(&memory.id))
.unwrap_or(0.5);
let momentum_ema = feedback_store
.and_then(|fs| {
let store = fs.read();
store.get_momentum(&memory.id).map(|m| m.ema_with_decay())
})
.unwrap_or(0.0);
let fused_score = weights.fuse_scores_full(
semantic_score,
entity_score,
tag_score,
importance,
momentum_ema,
access_count,
graph_strength,
);
let reason = if semantic_score > 0.0 && entity_score > 0.0 {
RelevanceReason::Combined
} else if entity_score > 0.0 {
RelevanceReason::EntityMatch
} else if semantic_score > 0.0 {
RelevanceReason::SemanticSimilarity
} else {
RelevanceReason::RecentImportant
};
let recency_boosted = self.apply_recency_boost(
fused_score,
memory.created_at,
config.recency_boost_hours,
config.recency_boost_multiplier,
);
let final_score = if entity_score > 0.0 {
(recency_boosted * config.graph_boost_multiplier).min(1.0)
} else {
recency_boosted
};
Some(SurfacedMemory {
id: memory.id.0.to_string(),
content: memory.experience.content.clone(),
memory_type: format!("{:?}", memory.experience.experience_type),
importance,
relevance_score: final_score,
relevance_reason: reason.clone(),
matched_entities,
semantic_similarity: if semantic_score > 0.0 {
Some(semantic_score)
} else {
None
},
created_at: memory.created_at,
tags: memory.experience.tags.clone(),
})
},
)
.collect();
results.sort_by(|a, b| b.relevance_score.total_cmp(&a.relevance_score));
const MIN_RELEVANCE_SCORE: f32 = 0.25;
results.retain(|r| r.relevance_score >= MIN_RELEVANCE_SCORE);
results.truncate(config.max_results);
debug.ranking_ms = ranking_start.elapsed().as_secs_f64() * 1000.0;
let total_latency = start.elapsed().as_secs_f64() * 1000.0;
let latency_target_met = total_latency < 30.0;
Ok(RelevanceResponse {
memories: results,
detected_entities,
latency_ms: total_latency,
latency_target_met,
debug: if cfg!(debug_assertions) {
Some(debug)
} else {
None
},
})
}
pub fn surface_relevant_with_momentum(
&self,
context: &str,
memory_system: &MemorySystem,
graph_memory: Option<&GraphMemory>,
config: &RelevanceConfig,
momentum_lookup: &HashMap<Uuid, f32>,
) -> Result<RelevanceResponse> {
let start = Instant::now();
let mut debug = RelevanceDebug {
ner_ms: 0.0,
entity_match_ms: 0.0,
semantic_search_ms: 0.0,
ranking_ms: 0.0,
memories_scanned: 0,
entity_matches: 0,
semantic_matches: 0,
};
let ner_start = Instant::now();
let detected_entities = if config.enable_entity_matching {
self.extract_entities(context)
} else {
Vec::new()
};
debug.ner_ms = ner_start.elapsed().as_secs_f64() * 1000.0;
let mut candidate_memories: HashMap<Uuid, (Memory, f32, f32, Vec<String>)> = HashMap::new();
if config.enable_entity_matching && !detected_entities.is_empty() {
let entity_start = Instant::now();
let entity_matches =
self.match_by_entities(&detected_entities, memory_system, graph_memory, config)?;
debug.entity_match_ms = entity_start.elapsed().as_secs_f64() * 1000.0;
debug.entity_matches = entity_matches.len();
for (memory, score, matched) in entity_matches {
let id = memory.id.0;
candidate_memories.insert(id, (memory, 0.0, score, matched));
}
}
if config.enable_semantic_matching {
let semantic_start = Instant::now();
let semantic_matches = self.match_by_semantic(context, memory_system, config)?;
debug.semantic_search_ms = semantic_start.elapsed().as_secs_f64() * 1000.0;
debug.semantic_matches = semantic_matches.len();
for (memory, score) in semantic_matches {
let id = memory.id.0;
if let Some((_, semantic_score, _entity_score, _matched)) =
candidate_memories.get_mut(&id)
{
*semantic_score = score;
} else {
candidate_memories.insert(id, (memory, score, 0.0, Vec::new()));
}
}
}
debug.memories_scanned = candidate_memories.len();
let ranking_start = Instant::now();
let weights = self.learned_weights.read().clone();
let mut results: Vec<SurfacedMemory> = candidate_memories
.into_iter()
.filter_map(
|(id, (memory, semantic_score, entity_score, matched_entities))| {
let importance = memory.importance();
if importance < config.min_importance {
return None;
}
if !config.memory_types.is_empty() {
let mem_type = format!("{:?}", memory.experience.experience_type);
if !config
.memory_types
.iter()
.any(|t| t.eq_ignore_ascii_case(&mem_type))
{
return None;
}
}
let tag_score = self.calculate_tag_score(context, &memory.experience.tags);
let momentum_ema = momentum_lookup.get(&id).copied().unwrap_or(0.0);
let access_count = memory.access_count();
let graph_strength = graph_memory
.and_then(|g| g.get_memory_hebbian_strength(&memory.id))
.unwrap_or(0.5);
let fused_score = weights.fuse_scores_full(
semantic_score,
entity_score,
tag_score,
importance,
momentum_ema,
access_count,
graph_strength,
);
let reason = if semantic_score > 0.0 && entity_score > 0.0 {
RelevanceReason::Combined
} else if entity_score > 0.0 {
RelevanceReason::EntityMatch
} else if semantic_score > 0.0 {
RelevanceReason::SemanticSimilarity
} else {
RelevanceReason::RecentImportant
};
let recency_boosted = self.apply_recency_boost(
fused_score,
memory.created_at,
config.recency_boost_hours,
config.recency_boost_multiplier,
);
let final_score = if entity_score > 0.0 {
(recency_boosted * config.graph_boost_multiplier).min(1.0)
} else {
recency_boosted
};
Some(SurfacedMemory {
id: memory.id.0.to_string(),
content: memory.experience.content.clone(),
memory_type: format!("{:?}", memory.experience.experience_type),
importance,
relevance_score: final_score,
relevance_reason: reason.clone(),
matched_entities,
semantic_similarity: if semantic_score > 0.0 {
Some(semantic_score)
} else {
None
},
created_at: memory.created_at,
tags: memory.experience.tags.clone(),
})
},
)
.collect();
results.sort_by(|a, b| b.relevance_score.total_cmp(&a.relevance_score));
const MIN_RELEVANCE_SCORE: f32 = 0.25;
results.retain(|r| r.relevance_score >= MIN_RELEVANCE_SCORE);
results.truncate(config.max_results);
debug.ranking_ms = ranking_start.elapsed().as_secs_f64() * 1000.0;
let total_latency = start.elapsed().as_secs_f64() * 1000.0;
let latency_target_met = total_latency < 30.0;
Ok(RelevanceResponse {
memories: results,
detected_entities,
latency_ms: total_latency,
latency_target_met,
debug: if cfg!(debug_assertions) {
Some(debug)
} else {
None
},
})
}
pub fn surface_relevant_with_ab_test(
&self,
context: &str,
user_id: &str,
memory_system: &MemorySystem,
graph_memory: Option<&GraphMemory>,
config: &RelevanceConfig,
ab_manager: &crate::ab_testing::ABTestManager,
) -> Result<(RelevanceResponse, Option<crate::ab_testing::ABTestVariant>)> {
let start = Instant::now();
let active_test = self.get_active_ab_test();
let (weights, variant) = if let Some(ref test_id) = active_test {
match ab_manager.get_weights_for_user(test_id, user_id) {
Ok(w) => {
let v = ab_manager.get_variant(test_id, user_id).ok();
(w, v)
}
Err(_) => {
(self.get_weights(), None)
}
}
} else {
(self.get_weights(), None)
};
let response = self.surface_relevant_inner(
context,
memory_system,
graph_memory,
config,
None,
Some(weights),
)?;
if let (Some(ref test_id), Some(ref _v)) = (&active_test, &variant) {
let latency_us = start.elapsed().as_micros() as u64;
let avg_score = if response.memories.is_empty() {
0.0
} else {
response
.memories
.iter()
.map(|m| m.relevance_score as f64)
.sum::<f64>()
/ response.memories.len() as f64
};
let _ = ab_manager.record_impression(test_id, user_id, avg_score, latency_us);
}
Ok((response, variant))
}
pub fn record_ab_click(
&self,
user_id: &str,
memory_id: Uuid,
ab_manager: &crate::ab_testing::ABTestManager,
) -> Result<()> {
if let Some(test_id) = self.get_active_ab_test() {
ab_manager
.record_click(&test_id, user_id, memory_id)
.map_err(|e| anyhow::anyhow!("Failed to record A/B click: {}", e))?;
}
Ok(())
}
pub fn record_ab_feedback(
&self,
user_id: &str,
positive: bool,
ab_manager: &crate::ab_testing::ABTestManager,
) -> Result<()> {
if let Some(test_id) = self.get_active_ab_test() {
ab_manager
.record_feedback(&test_id, user_id, positive)
.map_err(|e| anyhow::anyhow!("Failed to record A/B feedback: {}", e))?;
}
Ok(())
}
fn extract_entities(&self, context: &str) -> Vec<DetectedEntity> {
match self.ner.extract(context) {
Ok(entities) => entities
.into_iter()
.map(|e| DetectedEntity {
name: e.text,
entity_type: format!("{:?}", e.entity_type),
confidence: e.confidence,
})
.collect(),
Err(_) => Vec::new(),
}
}
fn match_by_entities(
&self,
entities: &[DetectedEntity],
memory_system: &MemorySystem,
graph_memory: Option<&GraphMemory>,
config: &RelevanceConfig,
) -> Result<Vec<(Memory, f32, Vec<String>)>> {
let entity_lookup: Vec<(String, &DetectedEntity, f32)> = entities
.iter()
.map(|e| {
let weight = self.entity_type_weight(&e.entity_type);
(e.name.to_lowercase(), e, weight)
})
.collect();
let max_candidates = config.max_results * 3;
let mut results: Vec<(Memory, f32, Vec<String>)> = Vec::with_capacity(max_candidates);
let mut found_ids: HashSet<Uuid> = HashSet::new();
{
let index = self.entity_index.read();
if !index.is_empty() {
let mut candidate_ids: HashMap<Uuid, (f32, Vec<String>)> = HashMap::new();
for (name_lower, entity, weight) in &entity_lookup {
if let Some(entry) = index.get(name_lower) {
for &memory_id in &entry.memory_ids {
let score = entity.confidence * weight;
candidate_ids
.entry(memory_id)
.and_modify(|(existing_score, matched)| {
*existing_score += score;
if !matched.contains(&entity.name) {
matched.push(entity.name.clone());
}
})
.or_insert((score, vec![entity.name.clone()]));
}
}
}
if !candidate_ids.is_empty() {
let mut sorted_candidates: Vec<_> = candidate_ids.into_iter().collect();
sorted_candidates.sort_by(|a, b| b.1 .0.total_cmp(&a.1 .0));
for (memory_id, (score, matched)) in
sorted_candidates.into_iter().take(max_candidates)
{
let normalized_score = (score / matched.len() as f32).min(1.0);
if normalized_score >= config.entity_threshold {
let mem_id = crate::memory::MemoryId(memory_id);
if let Ok(memory) = memory_system.get_memory(&mem_id) {
found_ids.insert(memory_id);
results.push((memory, normalized_score, matched));
}
}
}
}
if !results.is_empty() {
return Ok(results);
}
}
}
let all_memories = memory_system.get_all_memories()?;
if all_memories.is_empty() {
return Ok(results);
}
for shared_memory in &all_memories {
if results.len() >= max_candidates {
break;
}
if found_ids.contains(&shared_memory.id.0) {
continue;
}
let content_lower = shared_memory.experience.content.to_lowercase();
let mut matched: Vec<String> = Vec::new();
let mut match_score = 0.0f32;
for (name_lower, entity, weight) in &entity_lookup {
if contains_word(&content_lower, name_lower) {
matched.push(entity.name.clone());
match_score += entity.confidence * weight;
}
}
if !matched.is_empty() {
let normalized_score = (match_score / matched.len() as f32).min(1.0);
if normalized_score >= config.entity_threshold {
results.push(((**shared_memory).clone(), normalized_score, matched));
}
}
}
if results.len() >= config.max_results * 2 || graph_memory.is_none() {
return Ok(results);
}
let mut graph_results = results;
if let Some(graph) = graph_memory {
let found_ids: HashSet<Uuid> = graph_results.iter().map(|(m, _, _)| m.id.0).collect();
let max_graph_lookups = 5;
for (idx, entity) in entities.iter().enumerate() {
if idx >= max_graph_lookups {
break;
}
if let Ok(Some(entity_node)) = graph.find_entity_by_name(&entity.name) {
if let Ok(traversal) = graph.traverse_from_entity(&entity_node.uuid, 5) {
for traversed in &traversal.entities {
if let Ok(episodes) =
graph.get_episodes_by_entity(&traversed.entity.uuid)
{
for episode in episodes.iter().take(10) {
let memory_id = crate::memory::MemoryId(episode.uuid);
if found_ids.contains(&episode.uuid) {
continue;
}
let score = entity.confidence
* traversed.entity.salience
* traversed.decay_factor;
if score >= config.entity_threshold {
if let Ok(memory) = memory_system.get_memory(&memory_id) {
graph_results.push((
memory,
score,
vec![
entity.name.clone(),
traversed.entity.name.clone(),
],
));
}
}
}
}
}
}
}
}
}
Ok(graph_results)
}
fn entity_type_weight(&self, entity_type: &str) -> f32 {
match entity_type.to_lowercase().as_str() {
"person" => 1.0,
"organization" => 0.9,
"location" => 0.8,
"technology" => 0.85,
"product" => 0.9,
"event" => 0.7,
"date" => 0.5,
_ => 0.6,
}
}
fn match_by_semantic(
&self,
context: &str,
memory_system: &MemorySystem,
config: &RelevanceConfig,
) -> Result<Vec<(Memory, f32)>> {
let query = crate::memory::Query {
query_text: Some(context.to_string()),
max_results: config.max_results * 2, importance_threshold: Some(config.min_importance),
..Default::default()
};
let mut results: Vec<(Memory, f32)> = Vec::new();
match memory_system.recall(&query) {
Ok(shared_memories) => {
for (rank, shared_memory) in shared_memories.into_iter().enumerate() {
let memory = (*shared_memory).clone();
let score = shared_memory
.get_score()
.unwrap_or(1.0 / (rank as f32 + 1.0));
if score >= config.semantic_threshold {
results.push((memory, score));
}
}
}
Err(_) => {
let context_words: HashSet<&str> =
context.split_whitespace().filter(|w| w.len() > 3).collect();
let all_memories = memory_system.get_all_memories()?;
for shared_memory in all_memories {
let content_words: HashSet<&str> = shared_memory
.experience
.content
.split_whitespace()
.filter(|w| w.len() > 3)
.collect();
let overlap = context_words.intersection(&content_words).count();
if overlap > 0 {
let score = overlap as f32
/ (context_words.len() + content_words.len()) as f32
* 2.0;
if score >= config.semantic_threshold {
let memory = (*shared_memory).clone();
results.push((memory, score.min(1.0)));
}
}
}
}
}
Ok(results)
}
fn apply_recency_boost(
&self,
base_score: f32,
created_at: DateTime<Utc>,
boost_hours: u64,
multiplier: f32,
) -> f32 {
if boost_hours == 0 {
return base_score;
}
let now = Utc::now();
let age = now.signed_duration_since(created_at);
let age_hours = age.num_hours() as u64;
if age_hours <= boost_hours {
let decay = 1.0 - (age_hours as f32 / boost_hours as f32);
let boost = 1.0 + (multiplier - 1.0) * decay;
(base_score * boost).min(1.0)
} else {
base_score
}
}
pub fn refresh_entity_index(&self, graph_memory: &GraphMemory) -> Result<()> {
let mut index = self.entity_index.write();
index.clear();
let now = Utc::now();
let entities = graph_memory.get_all_entities()?;
for entity in entities {
let episodes = graph_memory.get_episodes_by_entity(&entity.uuid)?;
let episode_ids: HashSet<Uuid> = episodes.iter().map(|e| e.uuid).collect();
let name_lower = entity.name.to_lowercase();
index.insert(
name_lower,
EntityIndexEntry {
memory_ids: episode_ids,
last_updated: Some(now),
},
);
}
*self.entity_index_timestamp.write() = Some(now);
Ok(())
}
pub fn get_memories_for_entity(
&self,
entity_name: &str,
graph_memory: Option<&GraphMemory>,
) -> Option<HashSet<Uuid>> {
let name_lower = entity_name.to_lowercase();
{
let index = self.entity_index.read();
if let Some(entry) = index.get(&name_lower) {
return Some(entry.memory_ids.clone());
}
}
if let Some(graph) = graph_memory {
if let Ok(Some(entity)) = graph.find_entity_by_name(&name_lower) {
if let Ok(episodes) = graph.get_episodes_by_entity(&entity.uuid) {
let memory_ids: HashSet<Uuid> = episodes.iter().map(|e| e.uuid).collect();
let mut index = self.entity_index.write();
index.insert(
name_lower,
EntityIndexEntry {
memory_ids: memory_ids.clone(),
last_updated: Some(Utc::now()),
},
);
return Some(memory_ids);
}
}
}
None
}
pub fn entity_index_needs_refresh(&self, max_age_hours: i64) -> bool {
let timestamp = self.entity_index_timestamp.read();
match *timestamp {
None => true,
Some(ts) => {
let age = Utc::now().signed_duration_since(ts);
age.num_hours() > max_age_hours
}
}
}
pub fn entity_index_stats(&self) -> (usize, Option<DateTime<Utc>>) {
let index = self.entity_index.read();
let timestamp = *self.entity_index_timestamp.read();
(index.len(), timestamp)
}
pub fn clear_entity_index(&self) {
self.entity_index.write().clear();
*self.entity_index_timestamp.write() = None;
}
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct ContextMonitorHandshake {
pub user_id: String,
#[serde(default)]
pub config: Option<RelevanceConfig>,
#[serde(default = "default_debounce_ms")]
pub debounce_ms: u64,
}
fn default_debounce_ms() -> u64 {
100
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct ContextUpdate {
pub context: String,
#[serde(default)]
pub entities: Vec<String>,
#[serde(default)]
pub config: Option<RelevanceConfig>,
}
#[derive(Debug, Clone, Serialize, Deserialize)]
#[serde(tag = "type")]
pub enum ContextMonitorResponse {
#[serde(rename = "ack")]
Ack { timestamp: DateTime<Utc> },
#[serde(rename = "relevant")]
Relevant {
memories: Vec<SurfacedMemory>,
detected_entities: Vec<DetectedEntity>,
latency_ms: f64,
timestamp: DateTime<Utc>,
},
#[serde(rename = "none")]
None { timestamp: DateTime<Utc> },
#[serde(rename = "error")]
Error {
code: String,
message: String,
fatal: bool,
timestamp: DateTime<Utc>,
},
}
pub struct ContextMonitor {
engine: Arc<RelevanceEngine>,
default_config: RelevanceConfig,
debounce_ms: u64,
}
impl ContextMonitor {
pub fn new(engine: Arc<RelevanceEngine>, debounce_ms: u64) -> Self {
Self {
engine,
default_config: RelevanceConfig::default(),
debounce_ms,
}
}
pub fn debounce_ms(&self) -> u64 {
self.debounce_ms
}
pub fn set_config(&mut self, config: RelevanceConfig) {
self.default_config = config;
}
pub fn engine(&self) -> &Arc<RelevanceEngine> {
&self.engine
}
pub fn process_context(
&self,
context: &str,
memory_system: &MemorySystem,
graph_memory: Option<&GraphMemory>,
config: Option<&RelevanceConfig>,
) -> Result<Option<RelevanceResponse>> {
let cfg = config.unwrap_or(&self.default_config);
if context.len() < 10 {
return Ok(None);
}
let response =
self.engine
.surface_relevant(context, memory_system, graph_memory, cfg, None)?;
if response.memories.is_empty() {
Ok(None)
} else {
Ok(Some(response))
}
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_relevance_config_defaults() {
let config = RelevanceConfig::default();
assert_eq!(config.semantic_threshold, 0.45);
assert_eq!(config.entity_threshold, 0.5);
assert_eq!(config.max_results, 5);
assert!(config.enable_entity_matching);
assert!(config.enable_semantic_matching);
}
#[test]
fn test_recency_boost() {
let engine = RelevanceEngine::new(Arc::new(crate::embeddings::NeuralNer::new_fallback(
crate::embeddings::NerConfig::default(),
)));
let recent = Utc::now();
let boosted = engine.apply_recency_boost(0.5, recent, 24, 1.2);
assert!(boosted > 0.5);
let old = Utc::now() - chrono::Duration::hours(48);
let not_boosted = engine.apply_recency_boost(0.5, old, 24, 1.2);
assert!((not_boosted - 0.5).abs() < 0.001);
}
#[test]
fn test_entity_type_weight() {
let engine = RelevanceEngine::new(Arc::new(crate::embeddings::NeuralNer::new_fallback(
crate::embeddings::NerConfig::default(),
)));
assert_eq!(engine.entity_type_weight("Person"), 1.0);
assert_eq!(engine.entity_type_weight("organization"), 0.9);
assert!(engine.entity_type_weight("unknown") < 1.0);
}
#[test]
fn test_detected_entity_serialization() {
let entity = DetectedEntity {
name: "Rust".to_string(),
entity_type: "Technology".to_string(),
confidence: 0.95,
};
let json = serde_json::to_string(&entity).unwrap();
assert!(json.contains("Rust"));
assert!(json.contains("Technology"));
}
#[test]
fn test_learned_weights_default() {
let weights = LearnedWeights::default();
let sum = weights.semantic
+ weights.entity
+ weights.tag
+ weights.importance
+ weights.momentum
+ weights.access_count
+ weights.graph_strength;
assert!((sum - 1.0).abs() < 0.001);
assert_eq!(weights.semantic, DEFAULT_SEMANTIC_WEIGHT);
assert_eq!(weights.entity, DEFAULT_ENTITY_WEIGHT);
assert_eq!(weights.tag, DEFAULT_TAG_WEIGHT);
assert_eq!(weights.importance, DEFAULT_IMPORTANCE_WEIGHT);
assert_eq!(weights.momentum, DEFAULT_MOMENTUM_WEIGHT);
assert_eq!(weights.access_count, DEFAULT_ACCESS_COUNT_WEIGHT);
assert_eq!(weights.graph_strength, DEFAULT_GRAPH_STRENGTH_WEIGHT);
}
#[test]
fn test_learned_weights_normalize() {
let mut weights = LearnedWeights {
semantic: 0.5,
entity: 0.5,
tag: 0.5,
importance: 0.5,
momentum: 0.5,
access_count: 0.5,
graph_strength: 0.5, update_count: 0,
last_updated: None,
};
weights.normalize();
let sum = weights.semantic
+ weights.entity
+ weights.tag
+ weights.importance
+ weights.momentum
+ weights.access_count
+ weights.graph_strength;
assert!((sum - 1.0).abs() < 0.001);
assert!((weights.semantic - 1.0 / 7.0).abs() < 0.001);
}
#[test]
fn test_learned_weights_feedback_helpful() {
let mut weights = LearnedWeights::default();
let initial_semantic = weights.semantic;
let initial_entity = weights.entity;
weights.apply_feedback(true, true, false, true);
assert_eq!(weights.update_count, 1);
assert!(weights.last_updated.is_some());
let sum = weights.semantic
+ weights.entity
+ weights.tag
+ weights.importance
+ weights.momentum
+ weights.access_count
+ weights.graph_strength;
assert!((sum - 1.0).abs() < 0.001);
let new_se = weights.semantic + weights.entity;
let old_se = initial_semantic + initial_entity;
assert!(new_se >= old_se - 0.1); }
#[test]
fn test_learned_weights_feedback_not_helpful() {
let mut weights = LearnedWeights::default();
weights.apply_feedback(true, false, false, false);
let sum = weights.semantic
+ weights.entity
+ weights.tag
+ weights.importance
+ weights.momentum
+ weights.access_count
+ weights.graph_strength;
assert!((sum - 1.0).abs() < 0.001);
}
#[test]
fn test_calibrate_score() {
let high = calibrate_score(0.9);
assert!(high > 0.9);
let low = calibrate_score(0.1);
assert!(low < 0.1);
let mid = calibrate_score(SIGMOID_MIDPOINT);
assert!((mid - 0.5).abs() < 0.001);
}
#[test]
fn test_score_fusion() {
let weights = LearnedWeights::default();
let high = weights.fuse_scores_full(0.9, 0.9, 0.9, 0.9, 0.9, 16, 0.9);
assert!(high > 0.8, "high score was {}", high);
let low = weights.fuse_scores_full(0.1, 0.1, 0.1, 0.1, -0.9, 0, 0.1);
assert!(low < 0.3, "low score was {}", low);
let mixed = weights.fuse_scores_full(0.9, 0.1, 0.5, 0.7, 0.0, 2, 0.5);
assert!(mixed > 0.2 && mixed < 0.8, "mixed score was {}", mixed);
let legacy = weights.fuse_scores(0.9, 0.9, 0.9, 0.9);
assert!(legacy > 0.5, "legacy score was {}", legacy); }
#[test]
fn test_tag_score_calculation() {
let engine = RelevanceEngine::new(Arc::new(crate::embeddings::NeuralNer::new_fallback(
crate::embeddings::NerConfig::default(),
)));
let score = engine.calculate_tag_score("I love Rust programming", &["rust".to_string()]);
assert_eq!(score, 1.0);
let score = engine
.calculate_tag_score("Learning Rust", &["rust".to_string(), "python".to_string()]);
assert_eq!(score, 0.5);
let score = engine.calculate_tag_score("Hello world", &["rust".to_string()]);
assert_eq!(score, 0.0);
let score = engine.calculate_tag_score("Test", &[]);
assert_eq!(score, 0.0);
}
#[test]
fn test_min_weight_enforcement() {
let mut weights = LearnedWeights {
semantic: 0.1,
entity: MIN_WEIGHT + 0.01, tag: 0.3, importance: 0.1,
momentum: 0.1,
access_count: 0.1,
graph_strength: 0.1,
update_count: 0,
last_updated: None,
};
weights.apply_feedback(false, true, false, false);
assert!(
weights.entity >= MIN_WEIGHT,
"entity {} < MIN_WEIGHT {}",
weights.entity,
MIN_WEIGHT
);
let sum = weights.semantic
+ weights.entity
+ weights.tag
+ weights.importance
+ weights.momentum
+ weights.access_count
+ weights.graph_strength;
assert!((sum - 1.0).abs() < 0.001);
}
#[test]
fn test_fuse_scores_nan_inf_guard() {
let weights = LearnedWeights::default();
let result = weights.fuse_scores_full(f32::NAN, 0.5, 0.5, 0.5, 0.0, 1, 0.5);
assert!(result.is_finite(), "NaN input should produce finite output");
let result = weights.fuse_scores_full(0.5, f32::INFINITY, 0.5, 0.5, 0.0, 1, 0.5);
assert!(result.is_finite(), "Inf input should produce finite output");
let result = weights.fuse_scores_full(0.5, 0.5, f32::NEG_INFINITY, 0.5, 0.0, 1, 0.5);
assert!(
result.is_finite(),
"-Inf input should produce finite output"
);
let result = weights.fuse_scores_full(0.8, 0.5, 0.3, 0.6, 0.2, 3, 0.4);
assert!(result.is_finite());
assert!(result > 0.0);
}
}