Merge commit 'd803bfe2b1fe7f5e219e50ac20d6801a0a58ac75' as 'vendor/ruvector'

This commit is contained in:
ruv
2026-02-28 14:39:40 -05:00
7854 changed files with 3522914 additions and 0 deletions
@@ -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, &params);
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, &centroids);
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, &centroids)
}
/// 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);
}
}