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
458 lines
15 KiB
Rust
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);
|
|
}
|
|
}
|