use async_trait::async_trait;
use std::collections::HashMap;
use std::fmt::Debug;
use tokio::sync::RwLock;
#[async_trait]
pub trait SearchFeedback: Send + Sync + Debug {
async fn record_click(&self, query: &str, item_url: &str);
async fn record_irrelevant(&self, query: &str, item_url: &str);
async fn get_url_weight(&self, url: &str) -> f32;
}
#[derive(Debug)]
pub struct MemorySearchFeedback {
click_counts: RwLock<HashMap<String, u64>>,
irrelevant_counts: RwLock<HashMap<String, u64>>,
total_feedback: RwLock<u64>,
}
impl MemorySearchFeedback {
pub fn new() -> Self {
Self {
click_counts: RwLock::new(HashMap::new()),
irrelevant_counts: RwLock::new(HashMap::new()),
total_feedback: RwLock::new(0),
}
}
pub fn click_stats(&self) -> HashMap<String, u64> {
self.click_counts.try_read().map(|counts| counts.clone()).unwrap_or_default()
}
}
impl Default for MemorySearchFeedback {
fn default() -> Self {
Self::new()
}
}
#[async_trait]
impl SearchFeedback for MemorySearchFeedback {
async fn record_click(&self, _query: &str, item_url: &str) {
*self.click_counts.write().await.entry(item_url.to_string()).or_default() += 1;
*self.total_feedback.write().await += 1;
}
async fn record_irrelevant(&self, _query: &str, item_url: &str) {
*self.irrelevant_counts.write().await.entry(item_url.to_string()).or_default() += 1;
*self.total_feedback.write().await += 1;
}
async fn get_url_weight(&self, url: &str) -> f32 {
let clicks = self.click_counts.read().await;
let irrelevants = self.irrelevant_counts.read().await;
let c = clicks.get(url).copied().unwrap_or(0) as f32;
let i = irrelevants.get(url).copied().unwrap_or(0) as f32;
if c + i == 0.0 {
return 1.0; }
let score = 1.0 + (1.0 + c).log2() - (1.0 + i).log2();
score.max(0.1).min(5.0)
}
}