Skip to main content

oxirs_embed/
adaptive_learning.rs

1//! Adaptive Learning System for Real-Time Embedding Enhancement
2//!
3//! This module implements an advanced adaptive learning system that continuously
4//! improves embedding quality through online learning, feedback mechanisms,
5//! and dynamic model adaptation based on usage patterns.
6
7use anyhow::Result;
8use chrono::{DateTime, Utc};
9use nalgebra::{DMatrix, DVector};
10use serde::{Deserialize, Serialize};
11use std::collections::{HashMap, VecDeque};
12use std::sync::{Arc, RwLock};
13use std::time::{Duration, Instant};
14use tokio::sync::mpsc;
15use tracing::{debug, info, warn};
16
17/// Configuration for adaptive learning system
18#[derive(Debug, Clone, Serialize, Deserialize)]
19pub struct AdaptiveLearningConfig {
20    /// Learning rate for adaptation
21    pub learning_rate: f64,
22    /// Size of the experience buffer
23    pub buffer_size: usize,
24    /// Minimum samples before adaptation
25    pub min_samples_for_adaptation: usize,
26    /// Maximum adaptation frequency (per second)
27    pub max_adaptation_frequency: f64,
28    /// Quality threshold for positive feedback
29    pub quality_threshold: f64,
30    /// Enable meta-learning adaptation
31    pub enable_meta_learning: bool,
32    /// Batch size for adaptation updates
33    pub adaptation_batch_size: usize,
34}
35
36impl Default for AdaptiveLearningConfig {
37    fn default() -> Self {
38        Self {
39            learning_rate: 0.001,
40            buffer_size: 10000,
41            min_samples_for_adaptation: 100,
42            max_adaptation_frequency: 1.0,
43            quality_threshold: 0.8,
44            enable_meta_learning: true,
45            adaptation_batch_size: 32,
46        }
47    }
48}
49
50/// Feedback signal for embedding quality
51#[derive(Debug, Clone, Serialize, Deserialize)]
52pub struct QualityFeedback {
53    /// Query that generated the embedding
54    pub query: String,
55    /// Generated embedding
56    pub embedding: Vec<f64>,
57    /// Quality score (0.0 to 1.0)
58    pub quality_score: f64,
59    /// Timestamp of feedback  
60    #[serde(with = "chrono::serde::ts_seconds")]
61    pub timestamp: DateTime<Utc>,
62    /// User-provided relevance score
63    pub relevance: Option<f64>,
64    /// Task context
65    pub task_context: Option<String>,
66}
67
68/// Experience sample for adaptive learning
69#[derive(Debug, Clone)]
70pub struct ExperienceSample {
71    /// Input query/text
72    pub input: String,
73    /// Target embedding (from positive feedback)
74    pub target: Vec<f64>,
75    /// Current embedding (what model produced)
76    pub current: Vec<f64>,
77    /// Quality improvement needed
78    pub improvement_target: f64,
79    /// Context information
80    pub context: HashMap<String, String>,
81}
82
83/// Adaptation strategy for different learning scenarios
84#[derive(Debug, Clone, Serialize, Deserialize)]
85pub enum AdaptationStrategy {
86    /// Gradient-based fine-tuning
87    GradientDescent { momentum: f64, weight_decay: f64 },
88    /// Evolutionary adaptation
89    Evolutionary {
90        mutation_rate: f64,
91        population_size: usize,
92    },
93    /// Meta-learning adaptation (MAML-style)
94    MetaLearning {
95        inner_steps: usize,
96        outer_learning_rate: f64,
97    },
98    /// Bayesian optimization
99    BayesianOptimization {
100        exploration_factor: f64,
101        kernel_bandwidth: f64,
102    },
103}
104
105impl Default for AdaptationStrategy {
106    fn default() -> Self {
107        Self::GradientDescent {
108            momentum: 0.9,
109            weight_decay: 0.0001,
110        }
111    }
112}
113
114/// Performance metrics for adaptation tracking
115#[derive(Debug, Clone, Serialize, Deserialize)]
116pub struct AdaptationMetrics {
117    /// Number of adaptations performed
118    pub adaptations_count: usize,
119    /// Average quality improvement per adaptation
120    pub avg_quality_improvement: f64,
121    /// Current adaptation rate (adaptations per minute)
122    pub adaptation_rate: f64,
123    /// Experience buffer utilization
124    pub buffer_utilization: f64,
125    /// Model performance drift detection
126    pub performance_drift: f64,
127    /// Last adaptation timestamp
128    #[serde(with = "chrono::serde::ts_seconds_option")]
129    pub last_adaptation: Option<DateTime<Utc>>,
130}
131
132impl Default for AdaptationMetrics {
133    fn default() -> Self {
134        Self {
135            adaptations_count: 0,
136            avg_quality_improvement: 0.0,
137            adaptation_rate: 0.0,
138            buffer_utilization: 0.0,
139            performance_drift: 0.0,
140            last_adaptation: None,
141        }
142    }
143}
144
145/// Adaptive learning system for continuous embedding improvement
146pub struct AdaptiveLearningSystem {
147    /// Configuration
148    config: AdaptiveLearningConfig,
149    /// Experience buffer for storing feedback
150    experience_buffer: Arc<RwLock<VecDeque<ExperienceSample>>>,
151    /// Quality feedback receiver
152    feedback_receiver: Arc<RwLock<Option<mpsc::UnboundedReceiver<QualityFeedback>>>>,
153    /// Quality feedback sender
154    feedback_sender: mpsc::UnboundedSender<QualityFeedback>,
155    /// Adaptation strategy
156    strategy: AdaptationStrategy,
157    /// Current metrics
158    metrics: Arc<RwLock<AdaptationMetrics>>,
159    /// Model parameters for adaptation
160    model_parameters: Arc<RwLock<HashMap<String, DMatrix<f64>>>>,
161    /// Learning state
162    learning_state: Arc<RwLock<LearningState>>,
163    /// Last embedding actually produced for a given query, used as the
164    /// "current model output" reference point when new feedback arrives for
165    /// that same query. Absent until a query has been seen at least once.
166    last_known_embedding: Arc<RwLock<HashMap<String, Vec<f64>>>>,
167}
168
169/// Internal learning state
170#[derive(Debug, Clone)]
171struct LearningState {
172    /// Momentum vectors for gradient descent
173    momentum: HashMap<String, DMatrix<f64>>,
174    /// Adaptation history
175    adaptation_history: VecDeque<AdaptationRecord>,
176    /// Current learning rate (adaptive)
177    current_learning_rate: f64,
178    /// Performance baseline
179    performance_baseline: f64,
180}
181
182/// Record of adaptation for analysis
183#[derive(Debug, Clone)]
184#[allow(dead_code)]
185struct AdaptationRecord {
186    /// Timestamp of adaptation
187    timestamp: DateTime<Utc>,
188    /// Quality before adaptation
189    quality_before: f64,
190    /// Quality after adaptation
191    quality_after: f64,
192    /// Number of samples used
193    samples_used: usize,
194    /// Strategy used
195    strategy: AdaptationStrategy,
196}
197
198impl AdaptiveLearningSystem {
199    /// Create new adaptive learning system
200    pub fn new(config: AdaptiveLearningConfig) -> Self {
201        let (sender, receiver) = mpsc::unbounded_channel();
202        let learning_rate = config.learning_rate;
203
204        Self {
205            config,
206            experience_buffer: Arc::new(RwLock::new(VecDeque::new())),
207            feedback_receiver: Arc::new(RwLock::new(Some(receiver))),
208            feedback_sender: sender,
209            strategy: AdaptationStrategy::default(),
210            metrics: Arc::new(RwLock::new(AdaptationMetrics::default())),
211            model_parameters: Arc::new(RwLock::new(HashMap::new())),
212            learning_state: Arc::new(RwLock::new(LearningState {
213                momentum: HashMap::new(),
214                adaptation_history: VecDeque::new(),
215                current_learning_rate: learning_rate,
216                performance_baseline: 0.5,
217            })),
218            last_known_embedding: Arc::new(RwLock::new(HashMap::new())),
219        }
220    }
221
222    /// Create system with custom strategy
223    pub fn with_strategy(config: AdaptiveLearningConfig, strategy: AdaptationStrategy) -> Self {
224        let mut system = Self::new(config);
225        system.strategy = strategy;
226        system
227    }
228
229    /// Submit quality feedback
230    pub fn submit_feedback(&self, feedback: QualityFeedback) -> Result<()> {
231        self.feedback_sender.send(feedback)?;
232        Ok(())
233    }
234
235    /// Start the adaptive learning process
236    pub async fn start_learning(&self) -> Result<()> {
237        let mut receiver = self
238            .feedback_receiver
239            .write()
240            .expect("lock poisoned")
241            .take()
242            .ok_or_else(|| anyhow::anyhow!("Learning already started"))?;
243
244        info!("Starting adaptive learning system");
245
246        // Spawn feedback processing task
247        let experience_buffer = Arc::clone(&self.experience_buffer);
248        let metrics = Arc::clone(&self.metrics);
249        let config = self.config.clone();
250        let last_known_embedding = Arc::clone(&self.last_known_embedding);
251
252        tokio::spawn(async move {
253            while let Some(feedback) = receiver.recv().await {
254                if let Err(e) = Self::process_feedback(
255                    feedback,
256                    &experience_buffer,
257                    &metrics,
258                    &config,
259                    &last_known_embedding,
260                )
261                .await
262                {
263                    warn!("Error processing feedback: {}", e);
264                }
265            }
266        });
267
268        // Spawn adaptation task
269        let buffer = Arc::clone(&self.experience_buffer);
270        let metrics = Arc::clone(&self.metrics);
271        let parameters = Arc::clone(&self.model_parameters);
272        let learning_state = Arc::clone(&self.learning_state);
273        let config = self.config.clone();
274        let strategy = self.strategy.clone();
275
276        tokio::spawn(async move {
277            let mut last_adaptation = Instant::now();
278
279            loop {
280                tokio::time::sleep(Duration::from_millis(100)).await;
281
282                // Check if we should perform adaptation
283                let should_adapt = {
284                    let buffer_guard = buffer.read().expect("lock poisoned");
285                    let _metrics_guard = metrics.read().expect("lock poisoned");
286
287                    buffer_guard.len() >= config.min_samples_for_adaptation
288                        && last_adaptation.elapsed().as_secs_f64()
289                            >= 1.0 / config.max_adaptation_frequency
290                };
291
292                if should_adapt {
293                    match Self::perform_adaptation(
294                        &buffer,
295                        &metrics,
296                        &parameters,
297                        &learning_state,
298                        &config,
299                        &strategy,
300                    )
301                    .await
302                    {
303                        Err(e) => {
304                            warn!("Error during adaptation: {}", e);
305                        }
306                        _ => {
307                            last_adaptation = Instant::now();
308                        }
309                    }
310                }
311            }
312        });
313
314        Ok(())
315    }
316
317    /// Process incoming feedback
318    async fn process_feedback(
319        feedback: QualityFeedback,
320        buffer: &Arc<RwLock<VecDeque<ExperienceSample>>>,
321        metrics: &Arc<RwLock<AdaptationMetrics>>,
322        config: &AdaptiveLearningConfig,
323        last_known_embedding: &Arc<RwLock<HashMap<String, Vec<f64>>>>,
324    ) -> Result<()> {
325        // Convert feedback to experience sample
326        if feedback.quality_score > config.quality_threshold {
327            // `current` is the model's actual last-known output for this exact
328            // query: the embedding recorded the previous time feedback for this
329            // query was processed. For a query seen for the first time there is
330            // no prior model output to compare against, so we fall back to the
331            // freshly submitted embedding (a genuine cold-start, not a bug: a
332            // fresh query trivially "conforms" to itself until proven otherwise).
333            let current = {
334                let known_guard = last_known_embedding.read().expect("lock poisoned");
335                known_guard
336                    .get(&feedback.query)
337                    .cloned()
338                    .unwrap_or_else(|| feedback.embedding.clone())
339            };
340
341            let sample = ExperienceSample {
342                input: feedback.query.clone(),
343                target: feedback.embedding.clone(),
344                current,
345                improvement_target: 1.0 - feedback.quality_score,
346                context: feedback
347                    .task_context
348                    .map(|ctx| [("task".to_string(), ctx)].into())
349                    .unwrap_or_default(),
350            };
351
352            // Record this feedback's embedding as the new "current" reference
353            // for this query so the next round of feedback measures real drift.
354            {
355                let mut known_guard = last_known_embedding.write().expect("lock poisoned");
356                known_guard.insert(feedback.query.clone(), feedback.embedding.clone());
357            }
358
359            // Add to buffer
360            {
361                let mut buffer_guard = buffer.write().expect("lock poisoned");
362                buffer_guard.push_back(sample);
363
364                // Maintain buffer size
365                while buffer_guard.len() > config.buffer_size {
366                    buffer_guard.pop_front();
367                }
368            }
369
370            // Update metrics
371            {
372                let mut metrics_guard = metrics.write().expect("lock poisoned");
373                let buffer_guard = buffer.read().expect("lock poisoned");
374                metrics_guard.buffer_utilization =
375                    buffer_guard.len() as f64 / config.buffer_size as f64;
376            }
377
378            debug!(
379                "Processed feedback with quality score: {}",
380                feedback.quality_score
381            );
382        }
383
384        Ok(())
385    }
386
387    /// Perform model adaptation
388    async fn perform_adaptation(
389        buffer: &Arc<RwLock<VecDeque<ExperienceSample>>>,
390        metrics: &Arc<RwLock<AdaptationMetrics>>,
391        parameters: &Arc<RwLock<HashMap<String, DMatrix<f64>>>>,
392        learning_state: &Arc<RwLock<LearningState>>,
393        config: &AdaptiveLearningConfig,
394        strategy: &AdaptationStrategy,
395    ) -> Result<()> {
396        let samples = {
397            let buffer_guard = buffer.read().expect("lock poisoned");
398            buffer_guard
399                .iter()
400                .take(config.adaptation_batch_size)
401                .cloned()
402                .collect::<Vec<_>>()
403        };
404
405        if samples.is_empty() {
406            return Ok(());
407        }
408
409        info!("Performing adaptation with {} samples", samples.len());
410
411        // Calculate quality before adaptation
412        let quality_before = Self::calculate_current_quality(&samples)?;
413
414        // Perform adaptation based on strategy
415        match strategy {
416            AdaptationStrategy::GradientDescent {
417                momentum,
418                weight_decay,
419            } => {
420                Self::gradient_descent_adaptation(
421                    &samples,
422                    parameters,
423                    learning_state,
424                    *momentum,
425                    *weight_decay,
426                    config.learning_rate,
427                )?;
428            }
429            AdaptationStrategy::MetaLearning {
430                inner_steps,
431                outer_learning_rate,
432            } => {
433                Self::meta_learning_adaptation(
434                    &samples,
435                    parameters,
436                    learning_state,
437                    *inner_steps,
438                    *outer_learning_rate,
439                )?;
440            }
441            AdaptationStrategy::Evolutionary {
442                mutation_rate,
443                population_size,
444            } => {
445                Self::evolutionary_adaptation(
446                    &samples,
447                    parameters,
448                    *mutation_rate,
449                    *population_size,
450                )?;
451            }
452            AdaptationStrategy::BayesianOptimization {
453                exploration_factor,
454                kernel_bandwidth,
455            } => {
456                Self::bayesian_optimization_adaptation(
457                    &samples,
458                    parameters,
459                    *exploration_factor,
460                    *kernel_bandwidth,
461                )?;
462            }
463        }
464
465        // Calculate quality after adaptation
466        let quality_after = Self::calculate_current_quality(&samples)?;
467
468        // Update metrics
469        {
470            let mut metrics_guard = metrics.write().expect("lock poisoned");
471            metrics_guard.adaptations_count += 1;
472            let improvement = quality_after - quality_before;
473            metrics_guard.avg_quality_improvement = (metrics_guard.avg_quality_improvement
474                * (metrics_guard.adaptations_count - 1) as f64
475                + improvement)
476                / metrics_guard.adaptations_count as f64;
477            metrics_guard.last_adaptation = Some(Utc::now());
478        }
479
480        // Update learning state
481        {
482            let mut state_guard = learning_state.write().expect("lock poisoned");
483            state_guard.adaptation_history.push_back(AdaptationRecord {
484                timestamp: Utc::now(),
485                quality_before,
486                quality_after,
487                samples_used: samples.len(),
488                strategy: strategy.clone(),
489            });
490
491            // Maintain history size
492            while state_guard.adaptation_history.len() > 1000 {
493                state_guard.adaptation_history.pop_front();
494            }
495
496            // Adaptive learning rate
497            if quality_after > quality_before {
498                state_guard.current_learning_rate *= 1.01; // Increase slightly
499            } else {
500                state_guard.current_learning_rate *= 0.95; // Decrease
501            }
502            state_guard.current_learning_rate = state_guard
503                .current_learning_rate
504                .max(config.learning_rate * 0.1)
505                .min(config.learning_rate * 10.0);
506        }
507
508        info!(
509            "Adaptation completed: quality improved by {:.4}",
510            quality_after - quality_before
511        );
512
513        Ok(())
514    }
515
516    /// Calculate current quality based on samples
517    fn calculate_current_quality(samples: &[ExperienceSample]) -> Result<f64> {
518        if samples.is_empty() {
519            return Ok(0.0);
520        }
521
522        let total_quality: f64 = samples
523            .iter()
524            .map(|sample| {
525                // Calculate similarity between current and target embeddings
526                let current = DVector::from_vec(sample.current.clone());
527                let target = DVector::from_vec(sample.target.clone());
528
529                if current.len() != target.len() {
530                    return 0.0;
531                }
532
533                // Cosine similarity
534                let dot_product = current.dot(&target);
535                let norm_current = current.norm();
536                let norm_target = target.norm();
537
538                if norm_current == 0.0 || norm_target == 0.0 {
539                    return 0.0;
540                }
541
542                (dot_product / (norm_current * norm_target)).max(0.0)
543            })
544            .sum();
545
546        Ok(total_quality / samples.len() as f64)
547    }
548
549    /// Gradient descent adaptation
550    fn gradient_descent_adaptation(
551        samples: &[ExperienceSample],
552        parameters: &Arc<RwLock<HashMap<String, DMatrix<f64>>>>,
553        learning_state: &Arc<RwLock<LearningState>>,
554        momentum: f64,
555        weight_decay: f64,
556        learning_rate: f64,
557    ) -> Result<()> {
558        // Simplified gradient descent implementation
559        // In a real implementation, this would compute gradients based on the loss
560        // between current embeddings and target embeddings
561
562        let mut params_guard = parameters.write().expect("lock poisoned");
563        let mut state_guard = learning_state.write().expect("lock poisoned");
564
565        for (param_name, param_matrix) in params_guard.iter_mut() {
566            // Compute pseudo-gradient (simplified for demonstration)
567            let gradient = Self::compute_gradient(samples, param_matrix)?;
568
569            // Update momentum
570            let momentum_entry = state_guard
571                .momentum
572                .entry(param_name.clone())
573                .or_insert_with(|| DMatrix::zeros(param_matrix.nrows(), param_matrix.ncols()));
574
575            *momentum_entry = momentum_entry.clone() * momentum + &gradient;
576
577            // Apply weight decay
578            let decay_term = param_matrix.clone() * weight_decay;
579
580            // Update parameters
581            *param_matrix -= &(momentum_entry.clone() * learning_rate + decay_term * learning_rate);
582        }
583
584        Ok(())
585    }
586
587    /// Compute gradient (simplified placeholder)
588    fn compute_gradient(
589        _samples: &[ExperienceSample],
590        param_matrix: &DMatrix<f64>,
591    ) -> Result<DMatrix<f64>> {
592        // Simplified gradient computation
593        // In practice, this would involve backpropagation through the embedding model
594        let mut gradient = DMatrix::zeros(param_matrix.nrows(), param_matrix.ncols());
595
596        // Add small random perturbations as a placeholder
597        for i in 0..gradient.nrows() {
598            for j in 0..gradient.ncols() {
599                gradient[(i, j)] = ({
600                    use scirs2_core::random::{Random, RngExt};
601                    let mut random = Random::default();
602                    random.random::<f64>()
603                } - 0.5)
604                    * 0.001;
605            }
606        }
607
608        Ok(gradient)
609    }
610
611    /// Meta-learning adaptation (MAML-style)
612    fn meta_learning_adaptation(
613        samples: &[ExperienceSample],
614        parameters: &Arc<RwLock<HashMap<String, DMatrix<f64>>>>,
615        _learning_state: &Arc<RwLock<LearningState>>,
616        inner_steps: usize,
617        outer_learning_rate: f64,
618    ) -> Result<()> {
619        let mut params_guard = parameters.write().expect("lock poisoned");
620
621        // MAML inner loop
622        for _ in 0..inner_steps {
623            for (_, param_matrix) in params_guard.iter_mut() {
624                let gradient = Self::compute_gradient(samples, param_matrix)?;
625                *param_matrix -= &(gradient * outer_learning_rate);
626            }
627        }
628
629        Ok(())
630    }
631
632    /// Evolutionary adaptation
633    fn evolutionary_adaptation(
634        samples: &[ExperienceSample],
635        parameters: &Arc<RwLock<HashMap<String, DMatrix<f64>>>>,
636        mutation_rate: f64,
637        population_size: usize,
638    ) -> Result<()> {
639        let mut params_guard = parameters.write().expect("lock poisoned");
640
641        // Simple evolutionary strategy
642        for (_, param_matrix) in params_guard.iter_mut() {
643            let mut best_fitness = Self::evaluate_fitness(samples, param_matrix)?;
644            let mut best_params = param_matrix.clone();
645
646            // Generate population
647            for _ in 0..population_size {
648                let mut mutated = param_matrix.clone();
649
650                // Apply mutations
651                for i in 0..mutated.nrows() {
652                    for j in 0..mutated.ncols() {
653                        if {
654                            use scirs2_core::random::{Random, RngExt};
655                            let mut random = Random::default();
656                            random.random::<f64>()
657                        } < mutation_rate
658                        {
659                            mutated[(i, j)] += ({
660                                use scirs2_core::random::{Random, RngExt};
661                                let mut random = Random::default();
662                                random.random::<f64>()
663                            } - 0.5)
664                                * 0.01;
665                        }
666                    }
667                }
668
669                let fitness = Self::evaluate_fitness(samples, &mutated)?;
670                if fitness > best_fitness {
671                    best_fitness = fitness;
672                    best_params = mutated;
673                }
674            }
675
676            *param_matrix = best_params;
677        }
678
679        Ok(())
680    }
681
682    /// Evaluate fitness for evolutionary adaptation
683    fn evaluate_fitness(
684        _samples: &[ExperienceSample],
685        _param_matrix: &DMatrix<f64>,
686    ) -> Result<f64> {
687        // Simplified fitness evaluation
688        // In practice, this would evaluate how well the parameters perform on the samples
689        Ok({
690            use scirs2_core::random::{Random, RngExt};
691            let mut random = Random::default();
692            random.random::<f64>()
693        })
694    }
695
696    /// Bayesian optimization adaptation
697    fn bayesian_optimization_adaptation(
698        samples: &[ExperienceSample],
699        parameters: &Arc<RwLock<HashMap<String, DMatrix<f64>>>>,
700        exploration_factor: f64,
701        _kernel_bandwidth: f64,
702    ) -> Result<()> {
703        let mut params_guard = parameters.write().expect("lock poisoned");
704
705        // Simplified Bayesian optimization
706        for (_, param_matrix) in params_guard.iter_mut() {
707            let current_fitness = Self::evaluate_fitness(samples, param_matrix)?;
708
709            // Generate candidate solutions
710            let mut best_candidate = param_matrix.clone();
711            let mut best_acquisition = 0.0;
712
713            for _ in 0..10 {
714                let mut candidate = param_matrix.clone();
715
716                // Add exploration noise
717                for i in 0..candidate.nrows() {
718                    for j in 0..candidate.ncols() {
719                        candidate[(i, j)] += ({
720                            use scirs2_core::random::{Random, RngExt};
721                            let mut random = Random::default();
722                            random.random::<f64>()
723                        } - 0.5)
724                            * exploration_factor;
725                    }
726                }
727
728                let fitness = Self::evaluate_fitness(samples, &candidate)?;
729                let acquisition = fitness
730                    + exploration_factor * {
731                        use scirs2_core::random::{Random, RngExt};
732                        let mut random = Random::default();
733                        random.random::<f64>()
734                    };
735
736                if acquisition > best_acquisition {
737                    best_acquisition = acquisition;
738                    best_candidate = candidate;
739                }
740            }
741
742            // Update only if improvement is significant
743            if best_acquisition > current_fitness + 0.01 {
744                *param_matrix = best_candidate;
745            }
746        }
747
748        Ok(())
749    }
750
751    /// Get current metrics
752    pub fn get_metrics(&self) -> AdaptationMetrics {
753        self.metrics.read().expect("lock poisoned").clone()
754    }
755
756    /// Get feedback sender for external use
757    pub fn get_feedback_sender(&self) -> mpsc::UnboundedSender<QualityFeedback> {
758        self.feedback_sender.clone()
759    }
760
761    /// Update adaptation strategy
762    pub fn set_strategy(&mut self, strategy: AdaptationStrategy) {
763        self.strategy = strategy;
764    }
765
766    /// Reset learning state
767    pub fn reset_learning_state(&self) {
768        let mut state_guard = self.learning_state.write().expect("lock poisoned");
769        state_guard.momentum.clear();
770        state_guard.adaptation_history.clear();
771        state_guard.current_learning_rate = self.config.learning_rate;
772        state_guard.performance_baseline = 0.5;
773    }
774}
775
776#[cfg(test)]
777mod tests {
778    use super::*;
779    use tokio::time::{sleep, Duration};
780
781    #[tokio::test]
782    async fn test_adaptive_learning_system_creation() {
783        let config = AdaptiveLearningConfig::default();
784        let system = AdaptiveLearningSystem::new(config);
785
786        let metrics = system.get_metrics();
787        assert_eq!(metrics.adaptations_count, 0);
788        assert_eq!(metrics.avg_quality_improvement, 0.0);
789    }
790
791    #[tokio::test]
792    async fn test_feedback_submission() {
793        let config = AdaptiveLearningConfig::default();
794        let system = AdaptiveLearningSystem::new(config);
795
796        let feedback = QualityFeedback {
797            query: "test query".to_string(),
798            embedding: vec![0.1, 0.2, 0.3],
799            quality_score: 0.9,
800            timestamp: Utc::now(),
801            relevance: Some(0.8),
802            task_context: Some("similarity".to_string()),
803        };
804
805        assert!(system.submit_feedback(feedback).is_ok());
806    }
807
808    #[tokio::test]
809    async fn test_adaptive_learning_config_default() {
810        let config = AdaptiveLearningConfig::default();
811
812        assert_eq!(config.learning_rate, 0.001);
813        assert_eq!(config.buffer_size, 10000);
814        assert_eq!(config.min_samples_for_adaptation, 100);
815        assert_eq!(config.quality_threshold, 0.8);
816        assert!(config.enable_meta_learning);
817    }
818
819    #[tokio::test]
820    async fn test_adaptation_strategies() {
821        let config = AdaptiveLearningConfig::default();
822
823        // Test different strategies
824        let strategies = vec![
825            AdaptationStrategy::GradientDescent {
826                momentum: 0.9,
827                weight_decay: 0.0001,
828            },
829            AdaptationStrategy::MetaLearning {
830                inner_steps: 3,
831                outer_learning_rate: 0.01,
832            },
833            AdaptationStrategy::Evolutionary {
834                mutation_rate: 0.1,
835                population_size: 20,
836            },
837            AdaptationStrategy::BayesianOptimization {
838                exploration_factor: 0.1,
839                kernel_bandwidth: 1.0,
840            },
841        ];
842
843        for strategy in strategies {
844            let system = AdaptiveLearningSystem::with_strategy(config.clone(), strategy);
845            assert!(system.start_learning().await.is_ok());
846
847            // Give some time for initialization
848            sleep(Duration::from_millis(10)).await;
849        }
850    }
851
852    /// Regression test for the P1 finding: `current` must reflect the
853    /// model's actual last-known output for a query, not simply be copied
854    /// from the freshly submitted `target` embedding (which made the quality
855    /// signal driving adaptation always ~1.0 and meaningless).
856    #[tokio::test]
857    async fn test_process_feedback_current_reflects_prior_embedding_not_target() {
858        let buffer = Arc::new(RwLock::new(VecDeque::new()));
859        let metrics = Arc::new(RwLock::new(AdaptationMetrics::default()));
860        let config = AdaptiveLearningConfig::default();
861        let last_known_embedding = Arc::new(RwLock::new(HashMap::new()));
862
863        let first_feedback = QualityFeedback {
864            query: "q1".to_string(),
865            embedding: vec![1.0, 0.0, 0.0],
866            quality_score: 0.9,
867            timestamp: Utc::now(),
868            relevance: None,
869            task_context: None,
870        };
871        AdaptiveLearningSystem::process_feedback(
872            first_feedback,
873            &buffer,
874            &metrics,
875            &config,
876            &last_known_embedding,
877        )
878        .await
879        .expect("should succeed");
880
881        // First time seeing "q1": there is no prior output, so `current`
882        // legitimately equals `target` (a genuine cold start, not the bug).
883        {
884            let buf = buffer.read().expect("lock poisoned");
885            let sample = buf.back().expect("sample should exist");
886            assert_eq!(sample.current, sample.target);
887        }
888
889        let second_feedback = QualityFeedback {
890            query: "q1".to_string(),
891            embedding: vec![0.0, 1.0, 0.0],
892            quality_score: 0.9,
893            timestamp: Utc::now(),
894            relevance: None,
895            task_context: None,
896        };
897        AdaptiveLearningSystem::process_feedback(
898            second_feedback,
899            &buffer,
900            &metrics,
901            &config,
902            &last_known_embedding,
903        )
904        .await
905        .expect("should succeed");
906
907        // Second time seeing "q1": `current` must be the *previous* round's
908        // embedding, genuinely different from this round's `target`.
909        {
910            let buf = buffer.read().expect("lock poisoned");
911            let sample = buf.back().expect("sample should exist");
912            assert_eq!(sample.current, vec![1.0, 0.0, 0.0]);
913            assert_eq!(sample.target, vec![0.0, 1.0, 0.0]);
914            assert_ne!(sample.current, sample.target);
915        }
916    }
917
918    #[tokio::test]
919    async fn test_quality_calculation() {
920        let samples = vec![
921            ExperienceSample {
922                input: "test1".to_string(),
923                target: vec![1.0, 0.0, 0.0],
924                current: vec![0.9, 0.1, 0.0],
925                improvement_target: 0.1,
926                context: HashMap::new(),
927            },
928            ExperienceSample {
929                input: "test2".to_string(),
930                target: vec![0.0, 1.0, 0.0],
931                current: vec![0.0, 0.8, 0.2],
932                improvement_target: 0.2,
933                context: HashMap::new(),
934            },
935        ];
936
937        let quality =
938            AdaptiveLearningSystem::calculate_current_quality(&samples).expect("should succeed");
939        assert!(quality > 0.0 && quality <= 1.0);
940    }
941}