use crate::types::{Embedding, Pattern};
use std::collections::VecDeque;
pub struct ShortTermMemory {
window: VecDeque<PatternEntry>,
window_size: usize,
decay: f32,
}
#[allow(dead_code)]
struct PatternEntry {
pattern: Pattern,
attention_weight: f32,
insert_time: u64,
}
impl ShortTermMemory {
pub fn new(window_size: usize) -> Self {
Self {
window: VecDeque::with_capacity(window_size),
window_size,
decay: 0.99,
}
}
pub fn add(&mut self, pattern: Pattern) {
for entry in self.window.iter_mut() {
entry.attention_weight *= self.decay;
}
let entry = PatternEntry {
insert_time: pattern.created_at,
pattern,
attention_weight: 1.0,
};
self.window.push_back(entry);
while self.window.len() > self.window_size {
self.window.pop_front();
}
}
pub fn attention_score(&self, pattern: &Pattern) -> f32 {
if self.window.is_empty() {
return 0.0;
}
let mut total_score = 0.0;
let mut total_weight = 0.0;
for entry in self.window.iter() {
let similarity = pattern
.embedding
.cosine_similarity(&entry.pattern.embedding);
total_score += similarity * entry.attention_weight;
total_weight += entry.attention_weight;
}
if total_weight > 0.0 {
total_score / total_weight
} else {
0.0
}
}
pub fn search(&self, query: &Embedding, limit: usize) -> Vec<(Pattern, f32)> {
let mut results: Vec<_> = self
.window
.iter()
.map(|entry| {
let similarity = query.cosine_similarity(&entry.pattern.embedding);
(entry.pattern.clone(), similarity * entry.attention_weight)
})
.collect();
results.sort_by(|a, b| b.1.partial_cmp(&a.1).unwrap_or(std::cmp::Ordering::Equal));
results.truncate(limit);
results
}
pub fn max_similarity(&self, query: &Embedding) -> f32 {
self.window
.iter()
.map(|entry| query.cosine_similarity(&entry.pattern.embedding))
.fold(0.0_f32, f32::max)
}
pub fn len(&self) -> usize {
self.window.len()
}
pub fn is_empty(&self) -> bool {
self.window.is_empty()
}
pub fn clear(&mut self) {
self.window.clear();
}
pub fn set_decay(&mut self, decay: f32) {
self.decay = decay.clamp(0.0, 1.0);
}
pub fn get_active_patterns(&self, threshold: f32) -> Vec<&Pattern> {
self.window
.iter()
.filter(|e| e.attention_weight >= threshold)
.map(|e| &e.pattern)
.collect()
}
pub fn aggregate_embedding(&self) -> Option<Embedding> {
if self.window.is_empty() {
return None;
}
let dim = self.window.front()?.pattern.embedding.dim;
let mut aggregate = vec![0.0f32; dim];
let mut total_weight = 0.0f32;
for entry in self.window.iter() {
for (i, &v) in entry.pattern.embedding.vector.iter().enumerate() {
if i < dim {
aggregate[i] += v * entry.attention_weight;
}
}
total_weight += entry.attention_weight;
}
if total_weight > 0.0 {
for v in aggregate.iter_mut() {
*v /= total_weight;
}
}
Some(Embedding::new(aggregate))
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::types::pattern_id;
use std::collections::HashMap;
fn make_pattern(id: u8) -> Pattern {
let embedding = Embedding::new(vec![id as f32 / 255.0; 16]);
Pattern {
id: pattern_id(&[id]),
embedding,
metadata: HashMap::new(),
created_at: 1702656000000 + (id as u64 * 1000),
}
}
#[test]
fn test_add_and_retrieve() {
let mut stm = ShortTermMemory::new(10);
let p1 = make_pattern(1);
let p2 = make_pattern(2);
stm.add(p1.clone());
stm.add(p2.clone());
assert_eq!(stm.len(), 2);
}
#[test]
fn test_window_eviction() {
let mut stm = ShortTermMemory::new(3);
for i in 0..5 {
stm.add(make_pattern(i));
}
assert_eq!(stm.len(), 3);
}
#[test]
fn test_attention_decay() {
let mut stm = ShortTermMemory::new(10);
stm.set_decay(0.5);
let p1 = make_pattern(1);
let p2 = make_pattern(2);
stm.add(p1.clone());
stm.add(p2.clone());
let active = stm.get_active_patterns(0.6);
assert_eq!(active.len(), 1); }
#[test]
fn test_search() {
let mut stm = ShortTermMemory::new(10);
for i in 0..5 {
stm.add(make_pattern(i));
}
let query = Embedding::new(vec![2.0 / 255.0; 16]);
let results = stm.search(&query, 3);
assert_eq!(results.len(), 3);
assert!(results[0].1 >= results[1].1);
}
#[test]
fn test_aggregate_embedding() {
let mut stm = ShortTermMemory::new(10);
stm.add(make_pattern(10));
stm.add(make_pattern(20));
let aggregate = stm.aggregate_embedding().unwrap();
assert!(aggregate.vector[0] > 10.0 / 255.0);
assert!(aggregate.vector[0] < 20.0 / 255.0);
}
}