Files
Soundgen/crates/soundgen-feedback/src/lib.rs
T
Emil 3e2376ac26 Add feedback loop: SQLite-backed rating system + MCP tools + GUI Training tab
New crate: soundgen-feedback
- FeedbackDB with SQLite (rusqlite, bundled)
- CRUD: add, update_rating, get_all, get_unrated, top_rated, delete
- search_similar: keyword-based similarity for few-shot reference
- export_jsonl: export rated examples for fine-tuning
- 9 tests

MCP server: 5 tools (was 3)
- generate_batch: generate multiple sounds by name, store in DB
- get_reference_sounds: search DB for highly-rated similar sounds
- render_sound: now returns reference examples from DB in response
- --db flag to specify feedback database path

GUI: Training tab
- Open/create feedback DB (file dialog)
- Generate batch dialog (enter names, output dir)
- Sound list with star rating (1-5), feedback text, play button
- Export dataset as JSONL
- Stats display (total/rated/avg rating)

OpenCode config updated with --db feedback.db path
2026-06-21 22:40:07 +03:00

458 lines
15 KiB
Rust

//! 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<String>,
pub wav_path: Option<String>,
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<Self, String> {
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<Self, String> {
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<String, String> {
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<Vec<FeedbackEntry>, 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::<Result<Vec<_>, _>>()
.map_err(|e| format!("row: {}", e))
}
/// Get only unrated entries (rating = 0).
pub fn get_unrated(&self) -> Result<Vec<FeedbackEntry>, 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::<Result<Vec<_>, _>>()
.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<Vec<FeedbackEntry>, 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::<Result<Vec<_>, _>>()
.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<Vec<FeedbackEntry>, 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<DBStats, String> {
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<usize, String> {
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<FeedbackEntry> {
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<String> = row.get(4)?;
let wav_path: Option<String> = 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<String> {
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::<String>::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);
}
}