//! Feedback database — SQLite-backed storage for sound ratings. //! //! Stores generated sounds with their specs, WAV paths, and user ratings (1-5 stars). //! Provides similarity search by name keywords for few-shot reference examples. use std::path::Path; use rusqlite::{params, Connection}; use soundgen_fmt::SoundSpec; /// A feedback entry in the database. #[derive(Clone, Debug)] pub struct FeedbackEntry { pub id: String, pub name: String, pub spec: SoundSpec, pub rating: u8, pub feedback: Option, pub wav_path: Option, pub created_at: String, } /// Database statistics. #[derive(Clone, Debug, Default)] pub struct DBStats { pub total: usize, pub rated: usize, pub unrated: usize, pub avg_rating: f32, } /// SQLite-backed feedback database. pub struct FeedbackDB { conn: Connection, } impl FeedbackDB { /// Open or create a feedback database at the given path. pub fn open(path: &Path) -> Result { let conn = Connection::open(path).map_err(|e| format!("open DB: {}", e))?; conn.execute_batch( "CREATE TABLE IF NOT EXISTS feedback ( id TEXT PRIMARY KEY, name TEXT NOT NULL, spec_json TEXT NOT NULL, rating INTEGER DEFAULT 0, feedback TEXT, wav_path TEXT, created_at TEXT NOT NULL ); CREATE INDEX IF NOT EXISTS idx_name ON feedback(name); CREATE INDEX IF NOT EXISTS idx_rating ON feedback(rating);", ) .map_err(|e| format!("init DB: {}", e))?; Ok(Self { conn }) } /// Create an in-memory database (for tests). pub fn in_memory() -> Result { let conn = Connection::open_in_memory().map_err(|e| format!("open in-memory DB: {}", e))?; conn.execute_batch( "CREATE TABLE IF NOT EXISTS feedback ( id TEXT PRIMARY KEY, name TEXT NOT NULL, spec_json TEXT NOT NULL, rating INTEGER DEFAULT 0, feedback TEXT, wav_path TEXT, created_at TEXT NOT NULL );", ) .map_err(|e| format!("init DB: {}", e))?; Ok(Self { conn }) } /// Add a new sound entry. Returns the generated ID. pub fn add( &mut self, name: &str, spec: &SoundSpec, wav_path: Option<&str>, ) -> Result { let id = uuid::Uuid::new_v4().to_string(); let spec_json = serde_json::to_string(spec).map_err(|e| format!("serialize spec: {}", e))?; let now = chrono_now(); self.conn .execute( "INSERT INTO feedback (id, name, spec_json, rating, feedback, wav_path, created_at) VALUES (?1, ?2, ?3, 0, NULL, ?4, ?5)", params![id, name, spec_json, wav_path, now], ) .map_err(|e| format!("insert: {}", e))?; Ok(id) } /// Update the rating and feedback for an entry. pub fn update_rating( &mut self, id: &str, rating: u8, feedback: Option<&str>, ) -> Result<(), String> { let rating = rating.min(5) as i32; self.conn .execute( "UPDATE feedback SET rating = ?1, feedback = ?2 WHERE id = ?3", params![rating, feedback, id], ) .map_err(|e| format!("update: {}", e))?; Ok(()) } /// Get all entries, ordered by creation time (newest first). pub fn get_all(&self) -> Result, String> { let mut stmt = self .conn .prepare("SELECT id, name, spec_json, rating, feedback, wav_path, created_at FROM feedback ORDER BY created_at DESC") .map_err(|e| format!("prepare: {}", e))?; let rows = stmt .query_map([], row_to_entry) .map_err(|e| format!("query: {}", e))?; rows.collect::, _>>() .map_err(|e| format!("row: {}", e)) } /// Get only unrated entries (rating = 0). pub fn get_unrated(&self) -> Result, String> { let mut stmt = self .conn .prepare("SELECT id, name, spec_json, rating, feedback, wav_path, created_at FROM feedback WHERE rating = 0 ORDER BY created_at DESC") .map_err(|e| format!("prepare: {}", e))?; let rows = stmt .query_map([], row_to_entry) .map_err(|e| format!("query: {}", e))?; rows.collect::, _>>() .map_err(|e| format!("row: {}", e)) } /// Get the top-rated entries (rating >= min_rating), ordered by rating descending. pub fn top_rated(&self, limit: usize, min_rating: u8) -> Result, String> { let min_rating = min_rating.min(5) as i32; let mut stmt = self .conn .prepare("SELECT id, name, spec_json, rating, feedback, wav_path, created_at FROM feedback WHERE rating >= ?1 ORDER BY rating DESC, created_at DESC LIMIT ?2") .map_err(|e| format!("prepare: {}", e))?; let rows = stmt .query_map(params![min_rating, limit as i64], row_to_entry) .map_err(|e| format!("query: {}", e))?; rows.collect::, _>>() .map_err(|e| format!("row: {}", e)) } /// Search for entries with names similar to the given query. /// Uses keyword matching: splits both query and names on underscores/spaces, /// counts matching keywords, sorts by match count * rating. pub fn search_similar(&self, query: &str, limit: usize) -> Result, String> { let query_keywords = tokenize(query); if query_keywords.is_empty() { return self.top_rated(limit, 1); } let all = self.get_all()?; let mut scored: Vec<(FeedbackEntry, usize, u8)> = all .into_iter() .map(|entry| { let entry_keywords = tokenize(&entry.name); let match_count = query_keywords .iter() .filter(|qk| entry_keywords.iter().any(|ek| ek == *qk)) .count(); let rating = entry.rating; (entry, match_count, rating) }) .filter(|(_, count, _)| *count > 0) .collect(); // Sort by: match_count desc, then rating desc scored.sort_by(|a, b| b.1.cmp(&a.1).then_with(|| b.2.cmp(&a.2))); Ok(scored.into_iter().take(limit).map(|(e, _, _)| e).collect()) } /// Get database statistics. pub fn stats(&self) -> Result { let total: i64 = self .conn .query_row("SELECT COUNT(*) FROM feedback", [], |r| r.get(0)) .map_err(|e| format!("count: {}", e))?; let rated: i64 = self .conn .query_row("SELECT COUNT(*) FROM feedback WHERE rating > 0", [], |r| { r.get(0) }) .map_err(|e| format!("count: {}", e))?; let avg: f64 = self .conn .query_row( "SELECT AVG(rating) FROM feedback WHERE rating > 0", [], |r| r.get(0), ) .unwrap_or(0.0); Ok(DBStats { total: total as usize, rated: rated as usize, unrated: (total - rated) as usize, avg_rating: avg as f32, }) } /// Export rated entries (rating >= min_rating) as JSONL for fine-tuning. pub fn export_jsonl(&self, path: &Path, min_rating: u8) -> Result { let entries = self.top_rated(10000, min_rating)?; let mut content = String::new(); for entry in &entries { let spec_json = serde_json::to_string(&entry.spec).map_err(|e| format!("serialize: {}", e))?; let line = serde_json::json!({ "instruction": format!("Generate a SoundSpec for: {}", entry.name), "response": spec_json, "rating": entry.rating, }); content.push_str(&line.to_string()); content.push('\n'); } std::fs::write(path, content).map_err(|e| format!("write file: {}", e))?; Ok(entries.len()) } /// Delete an entry by ID. pub fn delete(&mut self, id: &str) -> Result<(), String> { self.conn .execute("DELETE FROM feedback WHERE id = ?1", params![id]) .map_err(|e| format!("delete: {}", e))?; Ok(()) } } fn row_to_entry(row: &rusqlite::Row) -> rusqlite::Result { let id: String = row.get(0)?; let name: String = row.get(1)?; let spec_json: String = row.get(2)?; let rating: i32 = row.get(3)?; let feedback: Option = row.get(4)?; let wav_path: Option = row.get(5)?; let created_at: String = row.get(6)?; let spec: SoundSpec = serde_json::from_str(&spec_json).unwrap_or_else(|_| SoundSpec::default()); Ok(FeedbackEntry { id, name, spec, rating: rating.clamp(0, 5) as u8, feedback, wav_path, created_at, }) } /// Tokenize a name into lowercase keywords. /// "missile_launch" → ["missile", "launch"] /// "Big Explosion" → ["big", "explosion"] fn tokenize(s: &str) -> Vec { s.split(|c: char| c == '_' || c == ' ' || c == '-') .map(|w| w.to_lowercase()) .filter(|w| !w.is_empty() && w.len() > 1) .collect() } fn chrono_now() -> String { // Simple timestamp without chrono dependency use std::time::{SystemTime, UNIX_EPOCH}; let secs = SystemTime::now() .duration_since(UNIX_EPOCH) .map(|d| d.as_secs()) .unwrap_or(0); format!("{}", secs) } #[cfg(test)] mod tests { use super::*; use soundgen_core::FrequencyAutomation; use soundgen_fmt::ChannelSpec; fn make_spec(name: &str) -> SoundSpec { SoundSpec { name: name.to_string(), duration: 0.1, sample_rate: 44100, channels: vec![ChannelSpec::Pulse { duty: 50, frequency: FrequencyAutomation::fixed(440.0), envelope: None, filter: None, volume: 0.5, pan: 0.0, }], } } #[test] fn test_open_and_add() { let mut db = FeedbackDB::in_memory().unwrap(); let spec = make_spec("test_sound"); let id = db.add("test_sound", &spec, Some("/tmp/test.wav")).unwrap(); assert!(!id.is_empty()); let all = db.get_all().unwrap(); assert_eq!(all.len(), 1); assert_eq!(all[0].name, "test_sound"); assert_eq!(all[0].rating, 0); } #[test] fn test_update_rating() { let mut db = FeedbackDB::in_memory().unwrap(); let spec = make_spec("test"); let id = db.add("test", &spec, None).unwrap(); db.update_rating(&id, 4, Some("good")).unwrap(); let all = db.get_all().unwrap(); assert_eq!(all[0].rating, 4); assert_eq!(all[0].feedback.as_deref(), Some("good")); } #[test] fn test_get_unrated() { let mut db = FeedbackDB::in_memory().unwrap(); let id1 = db.add("sound1", &make_spec("sound1"), None).unwrap(); let _id2 = db.add("sound2", &make_spec("sound2"), None).unwrap(); db.update_rating(&id1, 5, None).unwrap(); let unrated = db.get_unrated().unwrap(); assert_eq!(unrated.len(), 1); assert_eq!(unrated[0].name, "sound2"); } #[test] fn test_top_rated() { let mut db = FeedbackDB::in_memory().unwrap(); let id1 = db.add("best", &make_spec("best"), None).unwrap(); let id2 = db.add("ok", &make_spec("ok"), None).unwrap(); let id3 = db.add("bad", &make_spec("bad"), None).unwrap(); db.update_rating(&id1, 5, None).unwrap(); db.update_rating(&id2, 3, None).unwrap(); db.update_rating(&id3, 1, None).unwrap(); let top = db.top_rated(2, 3).unwrap(); assert_eq!(top.len(), 2); assert_eq!(top[0].name, "best"); assert_eq!(top[0].rating, 5); assert_eq!(top[1].name, "ok"); } #[test] fn test_search_similar() { let mut db = FeedbackDB::in_memory().unwrap(); let id1 = db .add("missile_launch", &make_spec("missile_launch"), None) .unwrap(); let id2 = db .add("rocket_launch", &make_spec("rocket_launch"), None) .unwrap(); let id3 = db .add("coin_collect", &make_spec("coin_collect"), None) .unwrap(); db.update_rating(&id1, 5, None).unwrap(); db.update_rating(&id2, 4, None).unwrap(); db.update_rating(&id3, 3, None).unwrap(); // Search "missile launch" should find missile_launch (2 keyword matches) first let results = db.search_similar("missile_launch", 3).unwrap(); assert!(!results.is_empty()); assert_eq!(results[0].name, "missile_launch"); // Search "launch" should find both missile_launch and rocket_launch let results = db.search_similar("launch", 3).unwrap(); assert_eq!(results.len(), 2); } #[test] fn test_stats() { let mut db = FeedbackDB::in_memory().unwrap(); let id1 = db.add("s1", &make_spec("s1"), None).unwrap(); let id2 = db.add("s2", &make_spec("s2"), None).unwrap(); db.update_rating(&id1, 4, None).unwrap(); db.update_rating(&id2, 2, None).unwrap(); let stats = db.stats().unwrap(); assert_eq!(stats.total, 2); assert_eq!(stats.rated, 2); assert_eq!(stats.unrated, 0); assert!((stats.avg_rating - 3.0).abs() < 0.1); } #[test] fn test_export_jsonl() { let mut db = FeedbackDB::in_memory().unwrap(); let id1 = db .add("good_sound", &make_spec("good_sound"), None) .unwrap(); db.update_rating(&id1, 5, Some("perfect")).unwrap(); let path = std::env::temp_dir().join("soundgen_test_export.jsonl"); let count = db.export_jsonl(&path, 4).unwrap(); assert_eq!(count, 1); let content = std::fs::read_to_string(&path).unwrap(); assert!(content.contains("good_sound")); assert!(content.contains("instruction")); let _ = std::fs::remove_file(&path); } #[test] fn test_tokenize() { assert_eq!(tokenize("missile_launch"), vec!["missile", "launch"]); assert_eq!(tokenize("Big Explosion"), vec!["big", "explosion"]); assert_eq!(tokenize("coin-collect"), vec!["coin", "collect"]); assert_eq!(tokenize("a"), Vec::::new()); } #[test] fn test_delete() { let mut db = FeedbackDB::in_memory().unwrap(); let id = db.add("temp", &make_spec("temp"), None).unwrap(); assert_eq!(db.get_all().unwrap().len(), 1); db.delete(&id).unwrap(); assert_eq!(db.get_all().unwrap().len(), 0); } }