mirror of
https://github.com/ruvnet/RuView
synced 2026-07-29 18:31:44 +00:00
Merge commit 'd803bfe2b1fe7f5e219e50ac20d6801a0a58ac75' as 'vendor/ruvector'
This commit is contained in:
@@ -0,0 +1,119 @@
|
||||
//! Self-Learning and ReasoningBank Module
|
||||
//!
|
||||
//! This module implements adaptive query optimization using trajectory tracking,
|
||||
//! pattern extraction, and learned parameter optimization.
|
||||
|
||||
pub mod operators;
|
||||
pub mod optimizer;
|
||||
pub mod patterns;
|
||||
pub mod reasoning_bank;
|
||||
pub mod trajectory;
|
||||
|
||||
pub use optimizer::{OptimizationTarget, SearchOptimizer, SearchParams};
|
||||
pub use patterns::{LearnedPattern, PatternExtractor};
|
||||
pub use reasoning_bank::ReasoningBank;
|
||||
pub use trajectory::{QueryTrajectory, TrajectoryTracker};
|
||||
|
||||
use dashmap::DashMap;
|
||||
use std::sync::Arc;
|
||||
|
||||
/// Global learning state manager
|
||||
pub struct LearningManager {
|
||||
/// Trajectory trackers per table
|
||||
trackers: DashMap<String, Arc<TrajectoryTracker>>,
|
||||
/// ReasoningBank instances per table
|
||||
reasoning_banks: DashMap<String, Arc<ReasoningBank>>,
|
||||
/// Search optimizers per table
|
||||
optimizers: DashMap<String, Arc<SearchOptimizer>>,
|
||||
}
|
||||
|
||||
impl LearningManager {
|
||||
/// Create a new learning manager
|
||||
pub fn new() -> Self {
|
||||
Self {
|
||||
trackers: DashMap::new(),
|
||||
reasoning_banks: DashMap::new(),
|
||||
optimizers: DashMap::new(),
|
||||
}
|
||||
}
|
||||
|
||||
/// Enable learning for a table
|
||||
pub fn enable_for_table(&self, table_name: &str, max_trajectories: usize) {
|
||||
let tracker = Arc::new(TrajectoryTracker::new(max_trajectories));
|
||||
let bank = Arc::new(ReasoningBank::new());
|
||||
let optimizer = Arc::new(SearchOptimizer::new(bank.clone()));
|
||||
|
||||
self.trackers.insert(table_name.to_string(), tracker);
|
||||
self.reasoning_banks.insert(table_name.to_string(), bank);
|
||||
self.optimizers.insert(table_name.to_string(), optimizer);
|
||||
}
|
||||
|
||||
/// Get tracker for a table
|
||||
pub fn get_tracker(&self, table_name: &str) -> Option<Arc<TrajectoryTracker>> {
|
||||
self.trackers.get(table_name).map(|r| r.value().clone())
|
||||
}
|
||||
|
||||
/// Get reasoning bank for a table
|
||||
pub fn get_reasoning_bank(&self, table_name: &str) -> Option<Arc<ReasoningBank>> {
|
||||
self.reasoning_banks
|
||||
.get(table_name)
|
||||
.map(|r| r.value().clone())
|
||||
}
|
||||
|
||||
/// Get optimizer for a table
|
||||
pub fn get_optimizer(&self, table_name: &str) -> Option<Arc<SearchOptimizer>> {
|
||||
self.optimizers.get(table_name).map(|r| r.value().clone())
|
||||
}
|
||||
|
||||
/// Extract and store patterns for a table
|
||||
pub fn extract_patterns(&self, table_name: &str, num_clusters: usize) -> Result<usize, String> {
|
||||
let tracker = self
|
||||
.get_tracker(table_name)
|
||||
.ok_or_else(|| format!("Learning not enabled for table: {}", table_name))?;
|
||||
let bank = self
|
||||
.get_reasoning_bank(table_name)
|
||||
.ok_or_else(|| format!("ReasoningBank not found for table: {}", table_name))?;
|
||||
|
||||
let trajectories = tracker.get_all();
|
||||
if trajectories.is_empty() {
|
||||
return Ok(0);
|
||||
}
|
||||
|
||||
let extractor = PatternExtractor::new(num_clusters);
|
||||
let patterns = extractor.extract_patterns(&trajectories);
|
||||
|
||||
let count = patterns.len();
|
||||
for pattern in patterns {
|
||||
bank.store(pattern);
|
||||
}
|
||||
|
||||
Ok(count)
|
||||
}
|
||||
}
|
||||
|
||||
impl Default for LearningManager {
|
||||
fn default() -> Self {
|
||||
Self::new()
|
||||
}
|
||||
}
|
||||
|
||||
lazy_static::lazy_static! {
|
||||
/// Global learning manager instance
|
||||
pub static ref LEARNING_MANAGER: LearningManager = LearningManager::new();
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::*;
|
||||
|
||||
#[test]
|
||||
fn test_learning_manager_lifecycle() {
|
||||
let manager = LearningManager::new();
|
||||
|
||||
manager.enable_for_table("test_table", 1000);
|
||||
|
||||
assert!(manager.get_tracker("test_table").is_some());
|
||||
assert!(manager.get_reasoning_bank("test_table").is_some());
|
||||
assert!(manager.get_optimizer("test_table").is_some());
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,533 @@
|
||||
//! PostgreSQL operator functions for self-learning
|
||||
|
||||
use pgrx::prelude::*;
|
||||
use pgrx::JsonB;
|
||||
use serde::{Deserialize, Serialize};
|
||||
|
||||
use super::optimizer::OptimizationTarget;
|
||||
use super::{QueryTrajectory, LEARNING_MANAGER};
|
||||
|
||||
/// Configuration for enabling learning
|
||||
#[derive(Debug, Serialize, Deserialize)]
|
||||
pub struct LearningConfig {
|
||||
/// Maximum number of trajectories to track
|
||||
#[serde(default = "default_max_trajectories")]
|
||||
pub max_trajectories: usize,
|
||||
/// Number of clusters for pattern extraction
|
||||
#[serde(default = "default_num_clusters")]
|
||||
pub num_clusters: usize,
|
||||
/// Auto-tune interval in seconds (0 = disabled)
|
||||
#[serde(default)]
|
||||
pub auto_tune_interval: u64,
|
||||
}
|
||||
|
||||
fn default_max_trajectories() -> usize {
|
||||
1000
|
||||
}
|
||||
fn default_num_clusters() -> usize {
|
||||
10
|
||||
}
|
||||
|
||||
impl Default for LearningConfig {
|
||||
fn default() -> Self {
|
||||
Self {
|
||||
max_trajectories: 1000,
|
||||
num_clusters: 10,
|
||||
auto_tune_interval: 0,
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
/// Enable learning for a table
|
||||
///
|
||||
/// # Examples
|
||||
///
|
||||
/// ```sql
|
||||
/// SELECT ruvector_enable_learning('my_table', '{"max_trajectories": 2000}'::jsonb);
|
||||
/// ```
|
||||
#[pg_extern]
|
||||
fn ruvector_enable_learning(
|
||||
table_name: &str,
|
||||
config: Option<JsonB>,
|
||||
) -> Result<String, Box<dyn std::error::Error + Send + Sync>> {
|
||||
let config: LearningConfig = match config {
|
||||
Some(jsonb) => serde_json::from_value(jsonb.0.clone())?,
|
||||
None => LearningConfig::default(),
|
||||
};
|
||||
|
||||
LEARNING_MANAGER.enable_for_table(table_name, config.max_trajectories);
|
||||
|
||||
Ok(format!(
|
||||
"Learning enabled for table '{}' with max_trajectories={}",
|
||||
table_name, config.max_trajectories
|
||||
))
|
||||
}
|
||||
|
||||
/// Record relevance feedback for a query
|
||||
///
|
||||
/// # Examples
|
||||
///
|
||||
/// ```sql
|
||||
/// SELECT ruvector_record_feedback(
|
||||
/// 'my_table',
|
||||
/// ARRAY[0.1, 0.2, 0.3],
|
||||
/// ARRAY[1, 2, 3]::bigint[],
|
||||
/// ARRAY[4, 5]::bigint[]
|
||||
/// );
|
||||
/// ```
|
||||
#[pg_extern]
|
||||
fn ruvector_record_feedback(
|
||||
table_name: &str,
|
||||
query_vector: Vec<f32>,
|
||||
relevant_ids: Vec<i64>,
|
||||
irrelevant_ids: Vec<i64>,
|
||||
) -> Result<String, Box<dyn std::error::Error + Send + Sync>> {
|
||||
let tracker = LEARNING_MANAGER
|
||||
.get_tracker(table_name)
|
||||
.ok_or_else(|| format!("Learning not enabled for table: {}", table_name))?;
|
||||
|
||||
// Find the most recent trajectory matching this query
|
||||
let mut recent = tracker.get_recent(10);
|
||||
|
||||
// Find matching trajectory (same query vector)
|
||||
if let Some(traj) = recent.iter_mut().find(|t| t.query_vector == query_vector) {
|
||||
traj.add_feedback(
|
||||
relevant_ids.iter().map(|&id| id as u64).collect(),
|
||||
irrelevant_ids.iter().map(|&id| id as u64).collect(),
|
||||
);
|
||||
|
||||
// Re-record the updated trajectory
|
||||
tracker.record(traj.clone());
|
||||
|
||||
Ok(format!(
|
||||
"Feedback recorded: {} relevant, {} irrelevant",
|
||||
relevant_ids.len(),
|
||||
irrelevant_ids.len()
|
||||
))
|
||||
} else {
|
||||
Err("No recent trajectory found matching query vector".into())
|
||||
}
|
||||
}
|
||||
|
||||
/// Get learning statistics for a table
|
||||
///
|
||||
/// # Examples
|
||||
///
|
||||
/// ```sql
|
||||
/// SELECT ruvector_learning_stats('my_table');
|
||||
/// ```
|
||||
#[pg_extern]
|
||||
fn ruvector_learning_stats(
|
||||
table_name: &str,
|
||||
) -> Result<JsonB, Box<dyn std::error::Error + Send + Sync>> {
|
||||
let tracker = LEARNING_MANAGER
|
||||
.get_tracker(table_name)
|
||||
.ok_or_else(|| format!("Learning not enabled for table: {}", table_name))?;
|
||||
|
||||
let bank = LEARNING_MANAGER
|
||||
.get_reasoning_bank(table_name)
|
||||
.ok_or_else(|| format!("ReasoningBank not found for table: {}", table_name))?;
|
||||
|
||||
let trajectory_stats = tracker.stats();
|
||||
let bank_stats = bank.stats();
|
||||
|
||||
let stats = serde_json::json!({
|
||||
"trajectories": {
|
||||
"total": trajectory_stats.total_trajectories,
|
||||
"with_feedback": trajectory_stats.trajectories_with_feedback,
|
||||
"avg_latency_us": trajectory_stats.avg_latency_us,
|
||||
"avg_precision": trajectory_stats.avg_precision,
|
||||
"avg_recall": trajectory_stats.avg_recall,
|
||||
},
|
||||
"patterns": {
|
||||
"total": bank_stats.total_patterns,
|
||||
"total_samples": bank_stats.total_samples,
|
||||
"avg_confidence": bank_stats.avg_confidence,
|
||||
"total_usage": bank_stats.total_usage,
|
||||
}
|
||||
});
|
||||
|
||||
Ok(JsonB(stats))
|
||||
}
|
||||
|
||||
/// Auto-tune search parameters for optimal performance
|
||||
///
|
||||
/// # Examples
|
||||
///
|
||||
/// ```sql
|
||||
/// SELECT ruvector_auto_tune(
|
||||
/// 'my_table',
|
||||
/// 'balanced',
|
||||
/// '[[0.1, 0.2, 0.3], [0.4, 0.5, 0.6]]'::jsonb
|
||||
/// );
|
||||
/// ```
|
||||
#[pg_extern]
|
||||
fn ruvector_auto_tune(
|
||||
table_name: &str,
|
||||
optimize_for: default!(&str, "'balanced'"),
|
||||
sample_queries: Option<JsonB>,
|
||||
) -> Result<JsonB, Box<dyn std::error::Error + Send + Sync>> {
|
||||
let optimizer = LEARNING_MANAGER
|
||||
.get_optimizer(table_name)
|
||||
.ok_or_else(|| format!("Learning not enabled for table: {}", table_name))?;
|
||||
|
||||
let target = match optimize_for {
|
||||
"speed" => OptimizationTarget::Speed,
|
||||
"accuracy" => OptimizationTarget::Accuracy,
|
||||
_ => OptimizationTarget::Balanced,
|
||||
};
|
||||
|
||||
// Extract patterns first
|
||||
let patterns_extracted = LEARNING_MANAGER.extract_patterns(table_name, 10)?;
|
||||
|
||||
let mut recommendations = Vec::new();
|
||||
|
||||
if let Some(JsonB(json_val)) = sample_queries {
|
||||
// Parse JSON array of arrays as Vec<Vec<f32>>
|
||||
if let Some(queries_array) = json_val.as_array() {
|
||||
for query_val in queries_array {
|
||||
if let Some(query_array) = query_val.as_array() {
|
||||
let query: Vec<f32> = query_array
|
||||
.iter()
|
||||
.filter_map(|v| v.as_f64().map(|f| f as f32))
|
||||
.collect();
|
||||
let params = optimizer.optimize_with_target(&query, target);
|
||||
recommendations.push(serde_json::json!({
|
||||
"ef_search": params.ef_search,
|
||||
"probes": params.probes,
|
||||
"confidence": params.confidence,
|
||||
}));
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
let result = serde_json::json!({
|
||||
"patterns_extracted": patterns_extracted,
|
||||
"optimize_for": optimize_for,
|
||||
"recommendations": recommendations,
|
||||
});
|
||||
|
||||
Ok(JsonB(result))
|
||||
}
|
||||
|
||||
/// Consolidate similar patterns to reduce memory usage
|
||||
///
|
||||
/// # Examples
|
||||
///
|
||||
/// ```sql
|
||||
/// SELECT ruvector_consolidate_patterns('my_table', 0.95);
|
||||
/// ```
|
||||
#[pg_extern]
|
||||
fn ruvector_consolidate_patterns(
|
||||
table_name: &str,
|
||||
similarity_threshold: default!(f64, 0.9),
|
||||
) -> Result<String, Box<dyn std::error::Error + Send + Sync>> {
|
||||
let bank = LEARNING_MANAGER
|
||||
.get_reasoning_bank(table_name)
|
||||
.ok_or_else(|| format!("Learning not enabled for table: {}", table_name))?;
|
||||
|
||||
let merged = bank.consolidate(similarity_threshold);
|
||||
|
||||
Ok(format!(
|
||||
"Consolidated {} similar patterns with threshold {}",
|
||||
merged, similarity_threshold
|
||||
))
|
||||
}
|
||||
|
||||
/// Prune low-quality patterns
|
||||
///
|
||||
/// # Examples
|
||||
///
|
||||
/// ```sql
|
||||
/// SELECT ruvector_prune_patterns('my_table', 5, 0.5);
|
||||
/// ```
|
||||
#[pg_extern]
|
||||
fn ruvector_prune_patterns(
|
||||
table_name: &str,
|
||||
min_usage: default!(i32, 5),
|
||||
min_confidence: default!(f64, 0.5),
|
||||
) -> Result<String, Box<dyn std::error::Error + Send + Sync>> {
|
||||
let bank = LEARNING_MANAGER
|
||||
.get_reasoning_bank(table_name)
|
||||
.ok_or_else(|| format!("Learning not enabled for table: {}", table_name))?;
|
||||
|
||||
let pruned = bank.prune(min_usage as usize, min_confidence);
|
||||
|
||||
Ok(format!(
|
||||
"Pruned {} patterns with min_usage={}, min_confidence={}",
|
||||
pruned, min_usage, min_confidence
|
||||
))
|
||||
}
|
||||
|
||||
/// Get optimized search parameters for a query
|
||||
///
|
||||
/// # Examples
|
||||
///
|
||||
/// ```sql
|
||||
/// SELECT ruvector_get_search_params('my_table', ARRAY[0.1, 0.2, 0.3]);
|
||||
/// ```
|
||||
#[pg_extern]
|
||||
fn ruvector_get_search_params(
|
||||
table_name: &str,
|
||||
query_vector: Vec<f32>,
|
||||
) -> Result<JsonB, Box<dyn std::error::Error + Send + Sync>> {
|
||||
let optimizer = LEARNING_MANAGER
|
||||
.get_optimizer(table_name)
|
||||
.ok_or_else(|| format!("Learning not enabled for table: {}", table_name))?;
|
||||
|
||||
let params = optimizer.optimize(&query_vector);
|
||||
|
||||
let result = serde_json::json!({
|
||||
"ef_search": params.ef_search,
|
||||
"probes": params.probes,
|
||||
"confidence": params.confidence,
|
||||
});
|
||||
|
||||
Ok(JsonB(result))
|
||||
}
|
||||
|
||||
/// Extract patterns from collected trajectories
|
||||
///
|
||||
/// # Examples
|
||||
///
|
||||
/// ```sql
|
||||
/// SELECT ruvector_extract_patterns('my_table', 10);
|
||||
/// ```
|
||||
#[pg_extern]
|
||||
fn ruvector_extract_patterns(
|
||||
table_name: &str,
|
||||
num_clusters: default!(i32, 10),
|
||||
) -> Result<String, Box<dyn std::error::Error + Send + Sync>> {
|
||||
let patterns_extracted =
|
||||
LEARNING_MANAGER.extract_patterns(table_name, num_clusters as usize)?;
|
||||
|
||||
Ok(format!(
|
||||
"Extracted {} patterns from trajectories using {} clusters",
|
||||
patterns_extracted, num_clusters
|
||||
))
|
||||
}
|
||||
|
||||
/// Record a query trajectory for learning
|
||||
///
|
||||
/// This is typically called internally by search functions, but can be used manually
|
||||
///
|
||||
/// # Examples
|
||||
///
|
||||
/// ```sql
|
||||
/// SELECT ruvector_record_trajectory(
|
||||
/// 'my_table',
|
||||
/// ARRAY[0.1, 0.2, 0.3],
|
||||
/// ARRAY[1, 2, 3]::bigint[],
|
||||
/// 1500,
|
||||
/// 50,
|
||||
/// 10
|
||||
/// );
|
||||
/// ```
|
||||
#[pg_extern]
|
||||
fn ruvector_record_trajectory(
|
||||
table_name: &str,
|
||||
query_vector: Vec<f32>,
|
||||
result_ids: Vec<i64>,
|
||||
latency_us: i64,
|
||||
ef_search: i32,
|
||||
probes: i32,
|
||||
) -> Result<String, Box<dyn std::error::Error + Send + Sync>> {
|
||||
let tracker = LEARNING_MANAGER
|
||||
.get_tracker(table_name)
|
||||
.ok_or_else(|| format!("Learning not enabled for table: {}", table_name))?;
|
||||
|
||||
let trajectory = QueryTrajectory::new(
|
||||
query_vector,
|
||||
result_ids.iter().map(|&id| id as u64).collect(),
|
||||
latency_us as u64,
|
||||
ef_search as usize,
|
||||
probes as usize,
|
||||
);
|
||||
|
||||
tracker.record(trajectory);
|
||||
|
||||
Ok(format!(
|
||||
"Trajectory recorded for {} results",
|
||||
result_ids.len()
|
||||
))
|
||||
}
|
||||
|
||||
/// Clear all learning data for a table
|
||||
///
|
||||
/// # Examples
|
||||
///
|
||||
/// ```sql
|
||||
/// SELECT ruvector_clear_learning('my_table');
|
||||
/// ```
|
||||
#[pg_extern]
|
||||
fn ruvector_clear_learning(
|
||||
table_name: &str,
|
||||
) -> Result<String, Box<dyn std::error::Error + Send + Sync>> {
|
||||
let bank = LEARNING_MANAGER
|
||||
.get_reasoning_bank(table_name)
|
||||
.ok_or_else(|| format!("Learning not enabled for table: {}", table_name))?;
|
||||
|
||||
bank.clear();
|
||||
|
||||
Ok(format!(
|
||||
"Cleared all learning data for table '{}'",
|
||||
table_name
|
||||
))
|
||||
}
|
||||
|
||||
#[cfg(feature = "pg_test")]
|
||||
#[pg_schema]
|
||||
mod tests {
|
||||
use super::*;
|
||||
|
||||
#[pg_test]
|
||||
fn test_enable_learning() {
|
||||
let result = ruvector_enable_learning("test_table", None);
|
||||
assert!(result.is_ok());
|
||||
}
|
||||
|
||||
#[pg_test]
|
||||
fn test_learning_stats_empty() {
|
||||
ruvector_enable_learning("test_stats", None).unwrap();
|
||||
let stats = ruvector_learning_stats("test_stats");
|
||||
assert!(stats.is_ok());
|
||||
}
|
||||
|
||||
#[pg_test]
|
||||
fn test_record_trajectory() {
|
||||
ruvector_enable_learning("test_trajectory", None).unwrap();
|
||||
|
||||
let result = ruvector_record_trajectory(
|
||||
"test_trajectory",
|
||||
vec![1.0, 2.0, 3.0],
|
||||
vec![1, 2, 3],
|
||||
1000,
|
||||
50,
|
||||
10,
|
||||
);
|
||||
|
||||
assert!(result.is_ok());
|
||||
}
|
||||
|
||||
#[pg_test]
|
||||
fn test_extract_patterns() {
|
||||
ruvector_enable_learning("test_patterns", None).unwrap();
|
||||
|
||||
// Record some trajectories
|
||||
for i in 0..20 {
|
||||
ruvector_record_trajectory(
|
||||
"test_patterns",
|
||||
vec![i as f32, (i * 2) as f32],
|
||||
vec![i, i + 1],
|
||||
1000 + i * 100,
|
||||
50,
|
||||
10,
|
||||
)
|
||||
.unwrap();
|
||||
}
|
||||
|
||||
let result = ruvector_extract_patterns("test_patterns", 5);
|
||||
assert!(result.is_ok());
|
||||
}
|
||||
|
||||
#[pg_test]
|
||||
fn test_auto_tune() {
|
||||
ruvector_enable_learning("test_autotune", None).unwrap();
|
||||
|
||||
// Record some trajectories
|
||||
for i in 0..10 {
|
||||
ruvector_record_trajectory(
|
||||
"test_autotune",
|
||||
vec![i as f32, (i * 2) as f32],
|
||||
vec![i],
|
||||
1000,
|
||||
50,
|
||||
10,
|
||||
)
|
||||
.unwrap();
|
||||
}
|
||||
|
||||
let result = ruvector_auto_tune("test_autotune", "balanced", None);
|
||||
|
||||
assert!(result.is_ok());
|
||||
}
|
||||
|
||||
#[pg_test]
|
||||
fn test_get_search_params() {
|
||||
ruvector_enable_learning("test_search_params", None).unwrap();
|
||||
|
||||
// Record and extract patterns first
|
||||
for i in 0..20 {
|
||||
ruvector_record_trajectory(
|
||||
"test_search_params",
|
||||
vec![i as f32, 0.0],
|
||||
vec![i],
|
||||
1000,
|
||||
50,
|
||||
10,
|
||||
)
|
||||
.unwrap();
|
||||
}
|
||||
|
||||
ruvector_extract_patterns("test_search_params", 3).unwrap();
|
||||
|
||||
let result = ruvector_get_search_params("test_search_params", vec![5.0, 0.0]);
|
||||
|
||||
assert!(result.is_ok());
|
||||
}
|
||||
|
||||
#[pg_test]
|
||||
fn test_consolidate_patterns() {
|
||||
ruvector_enable_learning("test_consolidate", None).unwrap();
|
||||
|
||||
// Record trajectories and extract patterns
|
||||
for i in 0..30 {
|
||||
ruvector_record_trajectory(
|
||||
"test_consolidate",
|
||||
vec![i as f32 / 10.0, 0.0],
|
||||
vec![i],
|
||||
1000,
|
||||
50,
|
||||
10,
|
||||
)
|
||||
.unwrap();
|
||||
}
|
||||
|
||||
ruvector_extract_patterns("test_consolidate", 10).unwrap();
|
||||
|
||||
let result = ruvector_consolidate_patterns("test_consolidate", 0.95);
|
||||
assert!(result.is_ok());
|
||||
}
|
||||
|
||||
#[pg_test]
|
||||
fn test_prune_patterns() {
|
||||
ruvector_enable_learning("test_prune", None).unwrap();
|
||||
|
||||
// Record trajectories and extract patterns
|
||||
for i in 0..20 {
|
||||
ruvector_record_trajectory("test_prune", vec![i as f32, 0.0], vec![i], 1000, 50, 10)
|
||||
.unwrap();
|
||||
}
|
||||
|
||||
ruvector_extract_patterns("test_prune", 5).unwrap();
|
||||
|
||||
let result = ruvector_prune_patterns("test_prune", 100, 0.9);
|
||||
assert!(result.is_ok());
|
||||
}
|
||||
|
||||
#[pg_test]
|
||||
fn test_clear_learning() {
|
||||
ruvector_enable_learning("test_clear", None).unwrap();
|
||||
|
||||
ruvector_record_trajectory("test_clear", vec![1.0, 2.0], vec![1], 1000, 50, 10).unwrap();
|
||||
|
||||
let result = ruvector_clear_learning("test_clear");
|
||||
assert!(result.is_ok());
|
||||
|
||||
let stats = ruvector_learning_stats("test_clear").unwrap();
|
||||
let stats_obj = stats.0.as_object().unwrap();
|
||||
let patterns = stats_obj.get("patterns").unwrap().as_object().unwrap();
|
||||
assert_eq!(patterns.get("total").unwrap().as_u64().unwrap(), 0);
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,332 @@
|
||||
//! Search parameter optimization using learned patterns
|
||||
|
||||
use super::reasoning_bank::ReasoningBank;
|
||||
use std::sync::Arc;
|
||||
|
||||
/// Search parameters for query execution
|
||||
#[derive(Debug, Clone, Copy)]
|
||||
pub struct SearchParams {
|
||||
pub ef_search: usize,
|
||||
pub probes: usize,
|
||||
pub confidence: f64,
|
||||
}
|
||||
|
||||
impl SearchParams {
|
||||
/// Create default search parameters
|
||||
pub fn default() -> Self {
|
||||
Self {
|
||||
ef_search: 50,
|
||||
probes: 10,
|
||||
confidence: 0.0,
|
||||
}
|
||||
}
|
||||
|
||||
/// Create with specific values
|
||||
pub fn new(ef_search: usize, probes: usize, confidence: f64) -> Self {
|
||||
Self {
|
||||
ef_search,
|
||||
probes,
|
||||
confidence,
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
/// Search optimizer using learned patterns
|
||||
pub struct SearchOptimizer {
|
||||
/// ReasoningBank for pattern lookup
|
||||
bank: Arc<ReasoningBank>,
|
||||
/// Number of patterns to consider
|
||||
k_patterns: usize,
|
||||
/// Minimum confidence threshold
|
||||
min_confidence: f64,
|
||||
}
|
||||
|
||||
impl SearchOptimizer {
|
||||
/// Create a new search optimizer
|
||||
pub fn new(bank: Arc<ReasoningBank>) -> Self {
|
||||
Self {
|
||||
bank,
|
||||
k_patterns: 5,
|
||||
min_confidence: 0.5,
|
||||
}
|
||||
}
|
||||
|
||||
/// Create with custom parameters
|
||||
pub fn with_params(bank: Arc<ReasoningBank>, k_patterns: usize, min_confidence: f64) -> Self {
|
||||
Self {
|
||||
bank,
|
||||
k_patterns,
|
||||
min_confidence,
|
||||
}
|
||||
}
|
||||
|
||||
/// Optimize search parameters for a query
|
||||
pub fn optimize(&self, query: &[f32]) -> SearchParams {
|
||||
// Lookup similar patterns
|
||||
let patterns = self.bank.lookup(query, self.k_patterns);
|
||||
|
||||
if patterns.is_empty() {
|
||||
return SearchParams::default();
|
||||
}
|
||||
|
||||
// Filter by confidence
|
||||
let valid_patterns: Vec<_> = patterns
|
||||
.iter()
|
||||
.filter(|(_, pattern, _)| pattern.confidence >= self.min_confidence)
|
||||
.collect();
|
||||
|
||||
if valid_patterns.is_empty() {
|
||||
return SearchParams::default();
|
||||
}
|
||||
|
||||
// Interpolate parameters based on similarity and confidence
|
||||
let mut total_weight = 0.0;
|
||||
let mut weighted_ef = 0.0;
|
||||
let mut weighted_probes = 0.0;
|
||||
let mut weighted_confidence = 0.0;
|
||||
|
||||
for (_, pattern, similarity) in valid_patterns.iter() {
|
||||
// Weight combines similarity and pattern confidence
|
||||
let weight = similarity * pattern.confidence;
|
||||
|
||||
weighted_ef += pattern.optimal_ef as f64 * weight;
|
||||
weighted_probes += pattern.optimal_probes as f64 * weight;
|
||||
weighted_confidence += pattern.confidence * weight;
|
||||
total_weight += weight;
|
||||
}
|
||||
|
||||
if total_weight == 0.0 {
|
||||
return SearchParams::default();
|
||||
}
|
||||
|
||||
SearchParams {
|
||||
ef_search: (weighted_ef / total_weight).round() as usize,
|
||||
probes: (weighted_probes / total_weight).round() as usize,
|
||||
confidence: weighted_confidence / total_weight,
|
||||
}
|
||||
}
|
||||
|
||||
/// Optimize with quality target (speed vs accuracy)
|
||||
pub fn optimize_with_target(&self, query: &[f32], target: OptimizationTarget) -> SearchParams {
|
||||
let mut params = self.optimize(query);
|
||||
|
||||
// Adjust based on target
|
||||
match target {
|
||||
OptimizationTarget::Speed => {
|
||||
// Reduce ef_search and probes for faster search
|
||||
params.ef_search = (params.ef_search as f64 * 0.7) as usize;
|
||||
params.probes = (params.probes as f64 * 0.7) as usize;
|
||||
}
|
||||
OptimizationTarget::Accuracy => {
|
||||
// Increase ef_search and probes for better accuracy
|
||||
params.ef_search = (params.ef_search as f64 * 1.3) as usize;
|
||||
params.probes = (params.probes as f64 * 1.3) as usize;
|
||||
}
|
||||
OptimizationTarget::Balanced => {
|
||||
// Use as-is
|
||||
}
|
||||
}
|
||||
|
||||
// Enforce minimum values
|
||||
params.ef_search = params.ef_search.max(10);
|
||||
params.probes = params.probes.max(1);
|
||||
|
||||
params
|
||||
}
|
||||
|
||||
/// Get recommendations for a query
|
||||
pub fn recommendations(&self, query: &[f32]) -> Vec<SearchRecommendation> {
|
||||
let patterns = self.bank.lookup(query, self.k_patterns);
|
||||
|
||||
patterns
|
||||
.iter()
|
||||
.filter(|(_, pattern, _)| pattern.confidence >= self.min_confidence)
|
||||
.map(|(id, pattern, similarity)| {
|
||||
let estimated_latency = pattern.avg_latency_us;
|
||||
let estimated_precision = pattern.avg_precision.unwrap_or(0.95);
|
||||
|
||||
SearchRecommendation {
|
||||
pattern_id: *id,
|
||||
ef_search: pattern.optimal_ef,
|
||||
probes: pattern.optimal_probes,
|
||||
similarity: *similarity,
|
||||
confidence: pattern.confidence,
|
||||
estimated_latency_us: estimated_latency,
|
||||
estimated_precision,
|
||||
}
|
||||
})
|
||||
.collect()
|
||||
}
|
||||
|
||||
/// Estimate query performance
|
||||
pub fn estimate_performance(
|
||||
&self,
|
||||
query: &[f32],
|
||||
params: &SearchParams,
|
||||
) -> PerformanceEstimate {
|
||||
let patterns = self.bank.lookup(query, self.k_patterns);
|
||||
|
||||
if patterns.is_empty() {
|
||||
return PerformanceEstimate::unknown();
|
||||
}
|
||||
|
||||
// Find patterns with similar parameters
|
||||
let similar_param_patterns: Vec<_> = patterns
|
||||
.iter()
|
||||
.filter(|(_, pattern, _)| {
|
||||
let ef_diff = (pattern.optimal_ef as i32 - params.ef_search as i32).abs();
|
||||
let probe_diff = (pattern.optimal_probes as i32 - params.probes as i32).abs();
|
||||
ef_diff < 20 && probe_diff < 5
|
||||
})
|
||||
.collect();
|
||||
|
||||
if similar_param_patterns.is_empty() {
|
||||
return PerformanceEstimate::low_confidence();
|
||||
}
|
||||
|
||||
// Weighted average of estimates
|
||||
let mut total_weight = 0.0;
|
||||
let mut weighted_latency = 0.0;
|
||||
let mut weighted_precision = 0.0;
|
||||
|
||||
for (_, pattern, similarity) in similar_param_patterns.iter() {
|
||||
let weight = similarity * pattern.confidence;
|
||||
weighted_latency += pattern.avg_latency_us * weight;
|
||||
if let Some(precision) = pattern.avg_precision {
|
||||
weighted_precision += precision * weight;
|
||||
}
|
||||
total_weight += weight;
|
||||
}
|
||||
|
||||
if total_weight == 0.0 {
|
||||
return PerformanceEstimate::low_confidence();
|
||||
}
|
||||
|
||||
PerformanceEstimate {
|
||||
estimated_latency_us: weighted_latency / total_weight,
|
||||
estimated_precision: Some(weighted_precision / total_weight),
|
||||
confidence: total_weight / similar_param_patterns.len() as f64,
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
/// Optimization target
|
||||
#[derive(Debug, Clone, Copy)]
|
||||
pub enum OptimizationTarget {
|
||||
Speed,
|
||||
Accuracy,
|
||||
Balanced,
|
||||
}
|
||||
|
||||
/// Search recommendation
|
||||
#[derive(Debug, Clone)]
|
||||
pub struct SearchRecommendation {
|
||||
pub pattern_id: usize,
|
||||
pub ef_search: usize,
|
||||
pub probes: usize,
|
||||
pub similarity: f64,
|
||||
pub confidence: f64,
|
||||
pub estimated_latency_us: f64,
|
||||
pub estimated_precision: f64,
|
||||
}
|
||||
|
||||
/// Performance estimate
|
||||
#[derive(Debug, Clone)]
|
||||
pub struct PerformanceEstimate {
|
||||
pub estimated_latency_us: f64,
|
||||
pub estimated_precision: Option<f64>,
|
||||
pub confidence: f64,
|
||||
}
|
||||
|
||||
impl PerformanceEstimate {
|
||||
fn unknown() -> Self {
|
||||
Self {
|
||||
estimated_latency_us: 0.0,
|
||||
estimated_precision: None,
|
||||
confidence: 0.0,
|
||||
}
|
||||
}
|
||||
|
||||
fn low_confidence() -> Self {
|
||||
Self {
|
||||
estimated_latency_us: 1000.0,
|
||||
estimated_precision: Some(0.9),
|
||||
confidence: 0.3,
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::*;
|
||||
use crate::learning::patterns::LearnedPattern;
|
||||
|
||||
fn create_test_bank() -> Arc<ReasoningBank> {
|
||||
let bank = Arc::new(ReasoningBank::new());
|
||||
|
||||
// Add test patterns
|
||||
let pattern1 =
|
||||
LearnedPattern::new(vec![1.0, 0.0, 0.0], 50, 10, 0.9, 100, 1000.0, Some(0.95));
|
||||
|
||||
let pattern2 =
|
||||
LearnedPattern::new(vec![0.0, 1.0, 0.0], 60, 15, 0.85, 80, 1500.0, Some(0.92));
|
||||
|
||||
bank.store(pattern1);
|
||||
bank.store(pattern2);
|
||||
|
||||
bank
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_optimize_basic() {
|
||||
let bank = create_test_bank();
|
||||
let optimizer = SearchOptimizer::new(bank);
|
||||
|
||||
let query = vec![0.9, 0.1, 0.0];
|
||||
let params = optimizer.optimize(&query);
|
||||
|
||||
assert!(params.ef_search > 0);
|
||||
assert!(params.probes > 0);
|
||||
assert!(params.confidence > 0.0);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_optimize_with_target() {
|
||||
let bank = create_test_bank();
|
||||
let optimizer = SearchOptimizer::new(bank);
|
||||
|
||||
let query = vec![1.0, 0.0, 0.0];
|
||||
|
||||
let speed_params = optimizer.optimize_with_target(&query, OptimizationTarget::Speed);
|
||||
let accuracy_params = optimizer.optimize_with_target(&query, OptimizationTarget::Accuracy);
|
||||
|
||||
assert!(speed_params.ef_search < accuracy_params.ef_search);
|
||||
assert!(speed_params.probes <= accuracy_params.probes);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_recommendations() {
|
||||
let bank = create_test_bank();
|
||||
let optimizer = SearchOptimizer::new(bank);
|
||||
|
||||
let query = vec![1.0, 0.0, 0.0];
|
||||
let recs = optimizer.recommendations(&query);
|
||||
|
||||
assert!(!recs.is_empty());
|
||||
assert!(recs[0].confidence >= 0.5);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_performance_estimate() {
|
||||
let bank = create_test_bank();
|
||||
let optimizer = SearchOptimizer::new(bank);
|
||||
|
||||
let query = vec![1.0, 0.0, 0.0];
|
||||
let params = SearchParams::new(50, 10, 0.9);
|
||||
|
||||
let estimate = optimizer.estimate_performance(&query, ¶ms);
|
||||
|
||||
assert!(estimate.estimated_latency_us > 0.0);
|
||||
assert!(estimate.confidence > 0.0);
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,361 @@
|
||||
//! Pattern extraction using k-means clustering
|
||||
|
||||
use super::trajectory::QueryTrajectory;
|
||||
|
||||
/// A learned pattern representing a cluster of similar queries
|
||||
#[derive(Debug, Clone)]
|
||||
pub struct LearnedPattern {
|
||||
/// Centroid vector of the pattern
|
||||
pub centroid: Vec<f32>,
|
||||
/// Optimal ef_search parameter for this pattern
|
||||
pub optimal_ef: usize,
|
||||
/// Optimal probes parameter for this pattern
|
||||
pub optimal_probes: usize,
|
||||
/// Confidence score (0.0 - 1.0)
|
||||
pub confidence: f64,
|
||||
/// Number of trajectories in this pattern
|
||||
pub sample_count: usize,
|
||||
/// Average latency for this pattern
|
||||
pub avg_latency_us: f64,
|
||||
/// Average precision (if feedback available)
|
||||
pub avg_precision: Option<f64>,
|
||||
}
|
||||
|
||||
impl LearnedPattern {
|
||||
/// Create a new pattern
|
||||
pub fn new(
|
||||
centroid: Vec<f32>,
|
||||
optimal_ef: usize,
|
||||
optimal_probes: usize,
|
||||
confidence: f64,
|
||||
sample_count: usize,
|
||||
avg_latency_us: f64,
|
||||
avg_precision: Option<f64>,
|
||||
) -> Self {
|
||||
Self {
|
||||
centroid,
|
||||
optimal_ef,
|
||||
optimal_probes,
|
||||
confidence,
|
||||
sample_count,
|
||||
avg_latency_us,
|
||||
avg_precision,
|
||||
}
|
||||
}
|
||||
|
||||
/// Calculate similarity to a query vector (cosine similarity)
|
||||
pub fn similarity(&self, query: &[f32]) -> f64 {
|
||||
if query.len() != self.centroid.len() {
|
||||
return 0.0;
|
||||
}
|
||||
|
||||
let dot: f32 = query.iter().zip(&self.centroid).map(|(a, b)| a * b).sum();
|
||||
let norm_q: f32 = query.iter().map(|x| x * x).sum::<f32>().sqrt();
|
||||
let norm_c: f32 = self.centroid.iter().map(|x| x * x).sum::<f32>().sqrt();
|
||||
|
||||
if norm_q == 0.0 || norm_c == 0.0 {
|
||||
return 0.0;
|
||||
}
|
||||
|
||||
(dot / (norm_q * norm_c)) as f64
|
||||
}
|
||||
}
|
||||
|
||||
/// Pattern extractor using k-means clustering
|
||||
pub struct PatternExtractor {
|
||||
/// Number of clusters
|
||||
k: usize,
|
||||
/// Maximum iterations for k-means
|
||||
max_iterations: usize,
|
||||
}
|
||||
|
||||
impl PatternExtractor {
|
||||
/// Create a new pattern extractor
|
||||
pub fn new(k: usize) -> Self {
|
||||
Self {
|
||||
k,
|
||||
max_iterations: 100,
|
||||
}
|
||||
}
|
||||
|
||||
/// Extract patterns from trajectories
|
||||
pub fn extract_patterns(&self, trajectories: &[QueryTrajectory]) -> Vec<LearnedPattern> {
|
||||
if trajectories.is_empty() || trajectories.len() < self.k {
|
||||
return Vec::new();
|
||||
}
|
||||
|
||||
let dim = trajectories[0].query_vector.len();
|
||||
|
||||
// Initialize centroids using k-means++
|
||||
let mut centroids = self.initialize_centroids(trajectories, dim);
|
||||
|
||||
// Run k-means
|
||||
let mut assignments = vec![0; trajectories.len()];
|
||||
|
||||
for _ in 0..self.max_iterations {
|
||||
let mut changed = false;
|
||||
|
||||
// Assignment step
|
||||
for (i, traj) in trajectories.iter().enumerate() {
|
||||
let closest = self.find_closest_centroid(&traj.query_vector, ¢roids);
|
||||
if assignments[i] != closest {
|
||||
assignments[i] = closest;
|
||||
changed = true;
|
||||
}
|
||||
}
|
||||
|
||||
if !changed {
|
||||
break;
|
||||
}
|
||||
|
||||
// Update step
|
||||
centroids = self.update_centroids(trajectories, &assignments, dim);
|
||||
}
|
||||
|
||||
// Create patterns from clusters
|
||||
self.create_patterns(trajectories, &assignments, ¢roids)
|
||||
}
|
||||
|
||||
/// Initialize centroids using k-means++
|
||||
fn initialize_centroids(
|
||||
&self,
|
||||
trajectories: &[QueryTrajectory],
|
||||
_default_ivfflat_probes: usize,
|
||||
) -> Vec<Vec<f32>> {
|
||||
let mut centroids = Vec::with_capacity(self.k);
|
||||
|
||||
// First centroid: random
|
||||
centroids.push(trajectories[0].query_vector.clone());
|
||||
|
||||
// Remaining centroids: weighted by distance
|
||||
for _ in 1..self.k {
|
||||
let mut distances = Vec::with_capacity(trajectories.len());
|
||||
|
||||
for traj in trajectories {
|
||||
let min_dist = centroids
|
||||
.iter()
|
||||
.map(|c| self.euclidean_distance(&traj.query_vector, c))
|
||||
.min_by(|a, b| a.partial_cmp(b).unwrap())
|
||||
.unwrap_or(0.0);
|
||||
distances.push(min_dist);
|
||||
}
|
||||
|
||||
// Select point with maximum distance
|
||||
let idx = distances
|
||||
.iter()
|
||||
.enumerate()
|
||||
.max_by(|(_, a), (_, b)| a.partial_cmp(b).unwrap())
|
||||
.map(|(i, _)| i)
|
||||
.unwrap_or(0);
|
||||
|
||||
centroids.push(trajectories[idx].query_vector.clone());
|
||||
}
|
||||
|
||||
centroids
|
||||
}
|
||||
|
||||
/// Find closest centroid index
|
||||
fn find_closest_centroid(&self, point: &[f32], centroids: &[Vec<f32>]) -> usize {
|
||||
centroids
|
||||
.iter()
|
||||
.enumerate()
|
||||
.map(|(i, c)| (i, self.euclidean_distance(point, c)))
|
||||
.min_by(|(_, a), (_, b)| a.partial_cmp(b).unwrap())
|
||||
.map(|(i, _)| i)
|
||||
.unwrap_or(0)
|
||||
}
|
||||
|
||||
/// Update centroids based on assignments
|
||||
fn update_centroids(
|
||||
&self,
|
||||
trajectories: &[QueryTrajectory],
|
||||
assignments: &[usize],
|
||||
dim: usize,
|
||||
) -> Vec<Vec<f32>> {
|
||||
let mut centroids = vec![vec![0.0; dim]; self.k];
|
||||
let mut counts = vec![0; self.k];
|
||||
|
||||
for (traj, &cluster) in trajectories.iter().zip(assignments) {
|
||||
for (i, &val) in traj.query_vector.iter().enumerate() {
|
||||
centroids[cluster][i] += val;
|
||||
}
|
||||
counts[cluster] += 1;
|
||||
}
|
||||
|
||||
for (centroid, &count) in centroids.iter_mut().zip(&counts) {
|
||||
if count > 0 {
|
||||
for val in centroid.iter_mut() {
|
||||
*val /= count as f32;
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
centroids
|
||||
}
|
||||
|
||||
/// Create patterns from clusters
|
||||
fn create_patterns(
|
||||
&self,
|
||||
trajectories: &[QueryTrajectory],
|
||||
assignments: &[usize],
|
||||
centroids: &[Vec<f32>],
|
||||
) -> Vec<LearnedPattern> {
|
||||
let mut patterns = Vec::new();
|
||||
|
||||
for cluster_id in 0..self.k {
|
||||
let cluster_trajs: Vec<&QueryTrajectory> = trajectories
|
||||
.iter()
|
||||
.zip(assignments)
|
||||
.filter(|(_, &a)| a == cluster_id)
|
||||
.map(|(t, _)| t)
|
||||
.collect();
|
||||
|
||||
if cluster_trajs.is_empty() {
|
||||
continue;
|
||||
}
|
||||
|
||||
// Calculate optimal parameters
|
||||
let optimal_ef = self.calculate_optimal_ef(&cluster_trajs);
|
||||
let optimal_probes = self.calculate_optimal_probes(&cluster_trajs);
|
||||
|
||||
// Calculate statistics
|
||||
let sample_count = cluster_trajs.len();
|
||||
let avg_latency = cluster_trajs.iter().map(|t| t.latency_us).sum::<u64>() as f64
|
||||
/ sample_count as f64;
|
||||
|
||||
let precisions: Vec<f64> = cluster_trajs.iter().filter_map(|t| t.precision()).collect();
|
||||
let avg_precision = if !precisions.is_empty() {
|
||||
Some(precisions.iter().sum::<f64>() / precisions.len() as f64)
|
||||
} else {
|
||||
None
|
||||
};
|
||||
|
||||
// Confidence based on sample count and consistency
|
||||
let confidence = self.calculate_confidence(&cluster_trajs);
|
||||
|
||||
patterns.push(LearnedPattern::new(
|
||||
centroids[cluster_id].clone(),
|
||||
optimal_ef,
|
||||
optimal_probes,
|
||||
confidence,
|
||||
sample_count,
|
||||
avg_latency,
|
||||
avg_precision,
|
||||
));
|
||||
}
|
||||
|
||||
patterns
|
||||
}
|
||||
|
||||
/// Calculate optimal ef_search for cluster
|
||||
fn calculate_optimal_ef(&self, trajectories: &[&QueryTrajectory]) -> usize {
|
||||
// Use median ef_search weighted by precision/latency trade-off
|
||||
let mut efs: Vec<_> = trajectories.iter().map(|t| t.ef_search).collect();
|
||||
efs.sort_unstable();
|
||||
|
||||
if efs.is_empty() {
|
||||
return 50; // Default
|
||||
}
|
||||
|
||||
efs[efs.len() / 2]
|
||||
}
|
||||
|
||||
/// Calculate optimal probes for cluster
|
||||
fn calculate_optimal_probes(&self, trajectories: &[&QueryTrajectory]) -> usize {
|
||||
let mut probes: Vec<_> = trajectories.iter().map(|t| t.probes).collect();
|
||||
probes.sort_unstable();
|
||||
|
||||
if probes.is_empty() {
|
||||
return 10; // Default
|
||||
}
|
||||
|
||||
probes[probes.len() / 2]
|
||||
}
|
||||
|
||||
/// Calculate confidence score
|
||||
fn calculate_confidence(&self, trajectories: &[&QueryTrajectory]) -> f64 {
|
||||
let n = trajectories.len() as f64;
|
||||
|
||||
// Base confidence on sample size
|
||||
let size_confidence = (n / 100.0).min(1.0);
|
||||
|
||||
// Consistency of parameters
|
||||
let ef_variance = self.calculate_variance(
|
||||
&trajectories
|
||||
.iter()
|
||||
.map(|t| t.ef_search as f64)
|
||||
.collect::<Vec<_>>(),
|
||||
);
|
||||
let consistency = 1.0 / (1.0 + ef_variance);
|
||||
|
||||
// Combined confidence
|
||||
(size_confidence * 0.7 + consistency * 0.3).min(1.0)
|
||||
}
|
||||
|
||||
/// Calculate variance
|
||||
fn calculate_variance(&self, values: &[f64]) -> f64 {
|
||||
if values.is_empty() {
|
||||
return 0.0;
|
||||
}
|
||||
|
||||
let mean = values.iter().sum::<f64>() / values.len() as f64;
|
||||
let variance = values.iter().map(|x| (x - mean).powi(2)).sum::<f64>() / values.len() as f64;
|
||||
|
||||
variance
|
||||
}
|
||||
|
||||
/// Euclidean distance between vectors
|
||||
fn euclidean_distance(&self, a: &[f32], b: &[f32]) -> f64 {
|
||||
a.iter()
|
||||
.zip(b)
|
||||
.map(|(x, y)| (x - y).powi(2))
|
||||
.sum::<f32>()
|
||||
.sqrt() as f64
|
||||
}
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::*;
|
||||
|
||||
#[test]
|
||||
fn test_pattern_similarity() {
|
||||
let pattern =
|
||||
LearnedPattern::new(vec![1.0, 0.0, 0.0], 50, 10, 0.9, 100, 1000.0, Some(0.95));
|
||||
|
||||
let query1 = vec![1.0, 0.0, 0.0]; // Same direction
|
||||
let query2 = vec![0.0, 1.0, 0.0]; // Perpendicular
|
||||
|
||||
assert!((pattern.similarity(&query1) - 1.0).abs() < 0.001);
|
||||
assert!((pattern.similarity(&query2) - 0.0).abs() < 0.001);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_pattern_extraction() {
|
||||
let trajectories = vec![
|
||||
QueryTrajectory::new(vec![1.0, 0.0], vec![1], 1000, 50, 10),
|
||||
QueryTrajectory::new(vec![1.1, 0.1], vec![1], 1100, 50, 10),
|
||||
QueryTrajectory::new(vec![0.0, 1.0], vec![2], 2000, 60, 15),
|
||||
QueryTrajectory::new(vec![0.1, 1.1], vec![2], 2100, 60, 15),
|
||||
];
|
||||
|
||||
let extractor = PatternExtractor::new(2);
|
||||
let patterns = extractor.extract_patterns(&trajectories);
|
||||
|
||||
assert_eq!(patterns.len(), 2);
|
||||
assert!(patterns.iter().all(|p| p.sample_count > 0));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_confidence_calculation() {
|
||||
let extractor = PatternExtractor::new(2);
|
||||
|
||||
// Consistent trajectories - create bindings first to avoid temporary value drop
|
||||
let traj1 = QueryTrajectory::new(vec![1.0], vec![1], 1000, 50, 10);
|
||||
let traj2 = QueryTrajectory::new(vec![1.0], vec![1], 1000, 50, 10);
|
||||
let trajs: Vec<&QueryTrajectory> = vec![&traj1, &traj2];
|
||||
|
||||
let confidence = extractor.calculate_confidence(&trajs);
|
||||
assert!(confidence > 0.0 && confidence <= 1.0);
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,337 @@
|
||||
//! ReasoningBank - Storage and retrieval of learned patterns
|
||||
|
||||
use super::patterns::LearnedPattern;
|
||||
use dashmap::DashMap;
|
||||
use std::sync::atomic::{AtomicUsize, Ordering};
|
||||
use std::time::SystemTime;
|
||||
|
||||
/// Pattern storage entry
|
||||
#[derive(Debug, Clone)]
|
||||
struct PatternEntry {
|
||||
pattern: LearnedPattern,
|
||||
usage_count: usize,
|
||||
last_used: SystemTime,
|
||||
}
|
||||
|
||||
/// ReasoningBank for storing and retrieving learned patterns
|
||||
pub struct ReasoningBank {
|
||||
/// Stored patterns indexed by ID
|
||||
patterns: DashMap<usize, PatternEntry>,
|
||||
/// Next pattern ID
|
||||
next_id: AtomicUsize,
|
||||
}
|
||||
|
||||
impl ReasoningBank {
|
||||
/// Create a new ReasoningBank
|
||||
pub fn new() -> Self {
|
||||
Self {
|
||||
patterns: DashMap::new(),
|
||||
next_id: AtomicUsize::new(0),
|
||||
}
|
||||
}
|
||||
|
||||
/// Store a new pattern
|
||||
pub fn store(&self, pattern: LearnedPattern) -> usize {
|
||||
let id = self.next_id.fetch_add(1, Ordering::SeqCst);
|
||||
|
||||
let entry = PatternEntry {
|
||||
pattern,
|
||||
usage_count: 0,
|
||||
last_used: SystemTime::now(),
|
||||
};
|
||||
|
||||
self.patterns.insert(id, entry);
|
||||
id
|
||||
}
|
||||
|
||||
/// Lookup k most similar patterns to a query
|
||||
pub fn lookup(&self, query: &[f32], k: usize) -> Vec<(usize, LearnedPattern, f64)> {
|
||||
let mut similarities: Vec<(usize, LearnedPattern, f64)> = self
|
||||
.patterns
|
||||
.iter()
|
||||
.map(|entry| {
|
||||
let id = *entry.key();
|
||||
let pattern = &entry.value().pattern;
|
||||
let similarity = pattern.similarity(query);
|
||||
(id, pattern.clone(), similarity)
|
||||
})
|
||||
.collect();
|
||||
|
||||
// Sort by similarity (descending) and confidence
|
||||
similarities.sort_by(|a, b| {
|
||||
let score_a = a.2 * a.1.confidence;
|
||||
let score_b = b.2 * b.1.confidence;
|
||||
score_b.partial_cmp(&score_a).unwrap()
|
||||
});
|
||||
|
||||
// Take top k
|
||||
similarities.truncate(k);
|
||||
|
||||
// Update usage statistics
|
||||
for (id, _, _) in &similarities {
|
||||
if let Some(mut entry) = self.patterns.get_mut(id) {
|
||||
entry.usage_count += 1;
|
||||
entry.last_used = SystemTime::now();
|
||||
}
|
||||
}
|
||||
|
||||
similarities
|
||||
}
|
||||
|
||||
/// Get a specific pattern by ID
|
||||
pub fn get(&self, id: usize) -> Option<LearnedPattern> {
|
||||
self.patterns.get_mut(&id).map(|mut entry| {
|
||||
entry.usage_count += 1;
|
||||
entry.last_used = SystemTime::now();
|
||||
entry.pattern.clone()
|
||||
})
|
||||
}
|
||||
|
||||
/// Consolidate similar patterns
|
||||
pub fn consolidate(&self, similarity_threshold: f64) -> usize {
|
||||
let patterns: Vec<(usize, LearnedPattern)> = self
|
||||
.patterns
|
||||
.iter()
|
||||
.map(|entry| (*entry.key(), entry.value().pattern.clone()))
|
||||
.collect();
|
||||
|
||||
if patterns.len() < 2 {
|
||||
return 0;
|
||||
}
|
||||
|
||||
let mut to_remove = Vec::new();
|
||||
let mut merged = 0;
|
||||
|
||||
for i in 0..patterns.len() {
|
||||
if to_remove.contains(&patterns[i].0) {
|
||||
continue;
|
||||
}
|
||||
|
||||
for j in (i + 1)..patterns.len() {
|
||||
if to_remove.contains(&patterns[j].0) {
|
||||
continue;
|
||||
}
|
||||
|
||||
let sim = patterns[i].1.similarity(&patterns[j].1.centroid);
|
||||
|
||||
if sim >= similarity_threshold {
|
||||
// Merge j into i
|
||||
if let Some(mut entry_i) = self.patterns.get_mut(&patterns[i].0) {
|
||||
if let Some(entry_j) = self.patterns.get(&patterns[j].0) {
|
||||
// Weighted merge based on sample counts
|
||||
let total_samples =
|
||||
entry_i.pattern.sample_count + entry_j.pattern.sample_count;
|
||||
let weight_i =
|
||||
entry_i.pattern.sample_count as f64 / total_samples as f64;
|
||||
let weight_j =
|
||||
entry_j.pattern.sample_count as f64 / total_samples as f64;
|
||||
|
||||
// Merge centroids
|
||||
for k in 0..entry_i.pattern.centroid.len() {
|
||||
entry_i.pattern.centroid[k] = (entry_i.pattern.centroid[k] as f64
|
||||
* weight_i
|
||||
+ entry_j.pattern.centroid[k] as f64 * weight_j)
|
||||
as f32;
|
||||
}
|
||||
|
||||
// Merge parameters (weighted average)
|
||||
entry_i.pattern.optimal_ef = (entry_i.pattern.optimal_ef as f64
|
||||
* weight_i
|
||||
+ entry_j.pattern.optimal_ef as f64 * weight_j)
|
||||
as usize;
|
||||
|
||||
entry_i.pattern.optimal_probes = (entry_i.pattern.optimal_probes as f64
|
||||
* weight_i
|
||||
+ entry_j.pattern.optimal_probes as f64 * weight_j)
|
||||
as usize;
|
||||
|
||||
// Update statistics
|
||||
entry_i.pattern.sample_count += entry_j.pattern.sample_count;
|
||||
entry_i.pattern.avg_latency_us = entry_i.pattern.avg_latency_us
|
||||
* weight_i
|
||||
+ entry_j.pattern.avg_latency_us * weight_j;
|
||||
|
||||
entry_i.pattern.confidence = (entry_i.pattern.confidence * weight_i
|
||||
+ entry_j.pattern.confidence * weight_j)
|
||||
.min(1.0);
|
||||
|
||||
entry_i.usage_count += entry_j.usage_count;
|
||||
}
|
||||
}
|
||||
|
||||
to_remove.push(patterns[j].0);
|
||||
merged += 1;
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// Remove merged patterns
|
||||
for id in to_remove {
|
||||
self.patterns.remove(&id);
|
||||
}
|
||||
|
||||
merged
|
||||
}
|
||||
|
||||
/// Prune low-quality patterns
|
||||
pub fn prune(&self, min_usage: usize, min_confidence: f64) -> usize {
|
||||
let to_remove: Vec<usize> = self
|
||||
.patterns
|
||||
.iter()
|
||||
.filter(|entry| {
|
||||
entry.value().usage_count < min_usage
|
||||
|| entry.value().pattern.confidence < min_confidence
|
||||
})
|
||||
.map(|entry| *entry.key())
|
||||
.collect();
|
||||
|
||||
let count = to_remove.len();
|
||||
for id in to_remove {
|
||||
self.patterns.remove(&id);
|
||||
}
|
||||
|
||||
count
|
||||
}
|
||||
|
||||
/// Get total number of patterns
|
||||
pub fn len(&self) -> usize {
|
||||
self.patterns.len()
|
||||
}
|
||||
|
||||
/// Check if bank is empty
|
||||
pub fn is_empty(&self) -> bool {
|
||||
self.patterns.is_empty()
|
||||
}
|
||||
|
||||
/// Get statistics
|
||||
pub fn stats(&self) -> BankStats {
|
||||
if self.patterns.is_empty() {
|
||||
return BankStats::default();
|
||||
}
|
||||
|
||||
let total = self.patterns.len();
|
||||
let total_samples: usize = self
|
||||
.patterns
|
||||
.iter()
|
||||
.map(|e| e.value().pattern.sample_count)
|
||||
.sum();
|
||||
|
||||
let avg_confidence: f64 = self
|
||||
.patterns
|
||||
.iter()
|
||||
.map(|e| e.value().pattern.confidence)
|
||||
.sum::<f64>()
|
||||
/ total as f64;
|
||||
|
||||
let total_usage: usize = self.patterns.iter().map(|e| e.value().usage_count).sum();
|
||||
|
||||
BankStats {
|
||||
total_patterns: total,
|
||||
total_samples,
|
||||
avg_confidence,
|
||||
total_usage,
|
||||
}
|
||||
}
|
||||
|
||||
/// Clear all patterns
|
||||
pub fn clear(&self) {
|
||||
self.patterns.clear();
|
||||
self.next_id.store(0, Ordering::SeqCst);
|
||||
}
|
||||
}
|
||||
|
||||
impl Default for ReasoningBank {
|
||||
fn default() -> Self {
|
||||
Self::new()
|
||||
}
|
||||
}
|
||||
|
||||
/// ReasoningBank statistics
|
||||
#[derive(Debug, Clone, Default)]
|
||||
pub struct BankStats {
|
||||
pub total_patterns: usize,
|
||||
pub total_samples: usize,
|
||||
pub avg_confidence: f64,
|
||||
pub total_usage: usize,
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::*;
|
||||
|
||||
fn create_test_pattern(centroid: Vec<f32>, ef: usize) -> LearnedPattern {
|
||||
LearnedPattern::new(centroid, ef, 10, 0.9, 100, 1000.0, Some(0.95))
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_store_and_lookup() {
|
||||
let bank = ReasoningBank::new();
|
||||
|
||||
let pattern1 = create_test_pattern(vec![1.0, 0.0, 0.0], 50);
|
||||
let pattern2 = create_test_pattern(vec![0.0, 1.0, 0.0], 60);
|
||||
|
||||
bank.store(pattern1);
|
||||
bank.store(pattern2);
|
||||
|
||||
assert_eq!(bank.len(), 2);
|
||||
|
||||
let query = vec![0.9, 0.1, 0.0];
|
||||
let results = bank.lookup(&query, 2);
|
||||
|
||||
assert_eq!(results.len(), 2);
|
||||
assert!(results[0].2 > results[1].2); // First result more similar
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_consolidate() {
|
||||
let bank = ReasoningBank::new();
|
||||
|
||||
// Store similar patterns
|
||||
let pattern1 = create_test_pattern(vec![1.0, 0.0], 50);
|
||||
let pattern2 = create_test_pattern(vec![0.99, 0.01], 50);
|
||||
let pattern3 = create_test_pattern(vec![0.0, 1.0], 60);
|
||||
|
||||
bank.store(pattern1);
|
||||
bank.store(pattern2);
|
||||
bank.store(pattern3);
|
||||
|
||||
assert_eq!(bank.len(), 3);
|
||||
|
||||
let merged = bank.consolidate(0.95);
|
||||
|
||||
assert!(merged > 0);
|
||||
assert!(bank.len() < 3);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_prune() {
|
||||
let bank = ReasoningBank::new();
|
||||
|
||||
let mut pattern_low_conf = create_test_pattern(vec![1.0, 0.0], 50);
|
||||
pattern_low_conf.confidence = 0.3;
|
||||
|
||||
bank.store(pattern_low_conf);
|
||||
bank.store(create_test_pattern(vec![0.0, 1.0], 60));
|
||||
|
||||
assert_eq!(bank.len(), 2);
|
||||
|
||||
let pruned = bank.prune(0, 0.5);
|
||||
|
||||
assert_eq!(pruned, 1);
|
||||
assert_eq!(bank.len(), 1);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_stats() {
|
||||
let bank = ReasoningBank::new();
|
||||
|
||||
bank.store(create_test_pattern(vec![1.0], 50));
|
||||
bank.store(create_test_pattern(vec![2.0], 60));
|
||||
|
||||
let stats = bank.stats();
|
||||
|
||||
assert_eq!(stats.total_patterns, 2);
|
||||
assert_eq!(stats.total_samples, 200);
|
||||
assert_eq!(stats.avg_confidence, 0.9);
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,287 @@
|
||||
//! Query trajectory tracking for learning query patterns
|
||||
|
||||
use std::sync::RwLock;
|
||||
use std::time::{Duration, SystemTime};
|
||||
|
||||
/// A single query trajectory record
|
||||
#[derive(Debug, Clone)]
|
||||
pub struct QueryTrajectory {
|
||||
/// Query vector
|
||||
pub query_vector: Vec<f32>,
|
||||
/// Result IDs
|
||||
pub result_ids: Vec<u64>,
|
||||
/// Query latency in microseconds
|
||||
pub latency_us: u64,
|
||||
/// Search parameters used
|
||||
pub ef_search: usize,
|
||||
pub probes: usize,
|
||||
/// Timestamp
|
||||
pub timestamp: SystemTime,
|
||||
/// Relevance feedback (if provided)
|
||||
pub relevant_ids: Vec<u64>,
|
||||
pub irrelevant_ids: Vec<u64>,
|
||||
}
|
||||
|
||||
impl QueryTrajectory {
|
||||
/// Create a new query trajectory
|
||||
pub fn new(
|
||||
query_vector: Vec<f32>,
|
||||
result_ids: Vec<u64>,
|
||||
latency_us: u64,
|
||||
ef_search: usize,
|
||||
probes: usize,
|
||||
) -> Self {
|
||||
Self {
|
||||
query_vector,
|
||||
result_ids,
|
||||
latency_us,
|
||||
ef_search,
|
||||
probes,
|
||||
timestamp: SystemTime::now(),
|
||||
relevant_ids: Vec::new(),
|
||||
irrelevant_ids: Vec::new(),
|
||||
}
|
||||
}
|
||||
|
||||
/// Add relevance feedback
|
||||
pub fn add_feedback(&mut self, relevant_ids: Vec<u64>, irrelevant_ids: Vec<u64>) {
|
||||
self.relevant_ids = relevant_ids;
|
||||
self.irrelevant_ids = irrelevant_ids;
|
||||
}
|
||||
|
||||
/// Calculate precision if feedback is available
|
||||
pub fn precision(&self) -> Option<f64> {
|
||||
if self.relevant_ids.is_empty() {
|
||||
return None;
|
||||
}
|
||||
|
||||
let relevant_retrieved = self
|
||||
.result_ids
|
||||
.iter()
|
||||
.filter(|id| self.relevant_ids.contains(id))
|
||||
.count();
|
||||
|
||||
Some(relevant_retrieved as f64 / self.result_ids.len() as f64)
|
||||
}
|
||||
|
||||
/// Calculate recall if feedback is available
|
||||
pub fn recall(&self) -> Option<f64> {
|
||||
if self.relevant_ids.is_empty() {
|
||||
return None;
|
||||
}
|
||||
|
||||
let relevant_retrieved = self
|
||||
.result_ids
|
||||
.iter()
|
||||
.filter(|id| self.relevant_ids.contains(id))
|
||||
.count();
|
||||
|
||||
Some(relevant_retrieved as f64 / self.relevant_ids.len() as f64)
|
||||
}
|
||||
}
|
||||
|
||||
/// Trajectory tracker with ring buffer
|
||||
pub struct TrajectoryTracker {
|
||||
/// Ring buffer of trajectories
|
||||
trajectories: RwLock<Vec<QueryTrajectory>>,
|
||||
/// Maximum number of trajectories to keep
|
||||
max_size: usize,
|
||||
/// Current write position
|
||||
write_pos: RwLock<usize>,
|
||||
}
|
||||
|
||||
impl TrajectoryTracker {
|
||||
/// Create a new trajectory tracker
|
||||
pub fn new(max_size: usize) -> Self {
|
||||
Self {
|
||||
trajectories: RwLock::new(Vec::with_capacity(max_size)),
|
||||
max_size,
|
||||
write_pos: RwLock::new(0),
|
||||
}
|
||||
}
|
||||
|
||||
/// Record a new trajectory
|
||||
pub fn record(&self, trajectory: QueryTrajectory) {
|
||||
let mut trajectories = self.trajectories.write().unwrap();
|
||||
let mut pos = self.write_pos.write().unwrap();
|
||||
|
||||
if trajectories.len() < self.max_size {
|
||||
trajectories.push(trajectory);
|
||||
} else {
|
||||
trajectories[*pos] = trajectory;
|
||||
}
|
||||
|
||||
*pos = (*pos + 1) % self.max_size;
|
||||
}
|
||||
|
||||
/// Get the most recent n trajectories
|
||||
pub fn get_recent(&self, n: usize) -> Vec<QueryTrajectory> {
|
||||
let trajectories = self.trajectories.read().unwrap();
|
||||
let count = trajectories.len().min(n);
|
||||
|
||||
if count == 0 {
|
||||
return Vec::new();
|
||||
}
|
||||
|
||||
let pos = *self.write_pos.read().unwrap();
|
||||
let mut result = Vec::with_capacity(count);
|
||||
|
||||
if trajectories.len() < self.max_size {
|
||||
// Not full yet, just take last n
|
||||
let start = trajectories.len().saturating_sub(count);
|
||||
result.extend_from_slice(&trajectories[start..]);
|
||||
} else {
|
||||
// Ring buffer is full, need to handle wrap-around
|
||||
for i in 0..count {
|
||||
let idx = (pos + self.max_size - count + i) % self.max_size;
|
||||
result.push(trajectories[idx].clone());
|
||||
}
|
||||
}
|
||||
|
||||
result
|
||||
}
|
||||
|
||||
/// Get all trajectories
|
||||
pub fn get_all(&self) -> Vec<QueryTrajectory> {
|
||||
self.trajectories.read().unwrap().clone()
|
||||
}
|
||||
|
||||
/// Get trajectories within a time window
|
||||
pub fn get_since(&self, duration: Duration) -> Vec<QueryTrajectory> {
|
||||
let trajectories = self.trajectories.read().unwrap();
|
||||
let cutoff = SystemTime::now() - duration;
|
||||
|
||||
trajectories
|
||||
.iter()
|
||||
.filter(|t| t.timestamp >= cutoff)
|
||||
.cloned()
|
||||
.collect()
|
||||
}
|
||||
|
||||
/// Get trajectories with feedback only
|
||||
pub fn get_with_feedback(&self) -> Vec<QueryTrajectory> {
|
||||
let trajectories = self.trajectories.read().unwrap();
|
||||
trajectories
|
||||
.iter()
|
||||
.filter(|t| !t.relevant_ids.is_empty())
|
||||
.cloned()
|
||||
.collect()
|
||||
}
|
||||
|
||||
/// Calculate average latency
|
||||
pub fn avg_latency(&self) -> Option<f64> {
|
||||
let trajectories = self.trajectories.read().unwrap();
|
||||
if trajectories.is_empty() {
|
||||
return None;
|
||||
}
|
||||
|
||||
let sum: u64 = trajectories.iter().map(|t| t.latency_us).sum();
|
||||
Some(sum as f64 / trajectories.len() as f64)
|
||||
}
|
||||
|
||||
/// Get statistics
|
||||
pub fn stats(&self) -> TrajectoryStats {
|
||||
let trajectories = self.trajectories.read().unwrap();
|
||||
|
||||
if trajectories.is_empty() {
|
||||
return TrajectoryStats::default();
|
||||
}
|
||||
|
||||
let total = trajectories.len();
|
||||
let with_feedback = trajectories
|
||||
.iter()
|
||||
.filter(|t| !t.relevant_ids.is_empty())
|
||||
.count();
|
||||
|
||||
let avg_latency =
|
||||
trajectories.iter().map(|t| t.latency_us).sum::<u64>() as f64 / total as f64;
|
||||
|
||||
let avg_precision = if with_feedback > 0 {
|
||||
trajectories
|
||||
.iter()
|
||||
.filter_map(|t| t.precision())
|
||||
.sum::<f64>()
|
||||
/ with_feedback as f64
|
||||
} else {
|
||||
0.0
|
||||
};
|
||||
|
||||
let avg_recall = if with_feedback > 0 {
|
||||
trajectories.iter().filter_map(|t| t.recall()).sum::<f64>() / with_feedback as f64
|
||||
} else {
|
||||
0.0
|
||||
};
|
||||
|
||||
TrajectoryStats {
|
||||
total_trajectories: total,
|
||||
trajectories_with_feedback: with_feedback,
|
||||
avg_latency_us: avg_latency,
|
||||
avg_precision,
|
||||
avg_recall,
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
/// Trajectory statistics
|
||||
#[derive(Debug, Clone, Default)]
|
||||
pub struct TrajectoryStats {
|
||||
pub total_trajectories: usize,
|
||||
pub trajectories_with_feedback: usize,
|
||||
pub avg_latency_us: f64,
|
||||
pub avg_precision: f64,
|
||||
pub avg_recall: f64,
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::*;
|
||||
|
||||
#[test]
|
||||
fn test_trajectory_creation() {
|
||||
let traj = QueryTrajectory::new(vec![1.0, 2.0, 3.0], vec![1, 2, 3], 1000, 50, 10);
|
||||
|
||||
assert_eq!(traj.query_vector, vec![1.0, 2.0, 3.0]);
|
||||
assert_eq!(traj.result_ids, vec![1, 2, 3]);
|
||||
assert_eq!(traj.latency_us, 1000);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_trajectory_feedback() {
|
||||
let mut traj = QueryTrajectory::new(vec![1.0, 2.0], vec![1, 2, 3, 4], 1000, 50, 10);
|
||||
|
||||
traj.add_feedback(vec![1, 2, 5], vec![3]);
|
||||
|
||||
assert_eq!(traj.precision(), Some(0.5)); // 2 out of 4 relevant
|
||||
assert_eq!(traj.recall(), Some(2.0 / 3.0)); // 2 out of 3 total relevant
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_tracker_ring_buffer() {
|
||||
let tracker = TrajectoryTracker::new(3);
|
||||
|
||||
// Add 5 trajectories
|
||||
for i in 0..5 {
|
||||
tracker.record(QueryTrajectory::new(vec![i as f32], vec![i], 1000, 50, 10));
|
||||
}
|
||||
|
||||
let all = tracker.get_all();
|
||||
assert_eq!(all.len(), 3); // Ring buffer size
|
||||
|
||||
// Should have trajectories 2, 3, 4 (last 3)
|
||||
let recent = tracker.get_recent(3);
|
||||
assert_eq!(recent.len(), 3);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_tracker_stats() {
|
||||
let tracker = TrajectoryTracker::new(10);
|
||||
|
||||
tracker.record(QueryTrajectory::new(vec![1.0], vec![1, 2], 1000, 50, 10));
|
||||
|
||||
tracker.record(QueryTrajectory::new(vec![2.0], vec![3, 4], 2000, 60, 15));
|
||||
|
||||
let stats = tracker.stats();
|
||||
assert_eq!(stats.total_trajectories, 2);
|
||||
assert_eq!(stats.avg_latency_us, 1500.0);
|
||||
}
|
||||
}
|
||||
Reference in New Issue
Block a user