Phase 4
This commit is contained in:
+351
-20
@@ -6,20 +6,25 @@
|
||||
//! else is still a typed `todo!()` stub mapped to its roadmap task.
|
||||
|
||||
use crate::audio::{AudioCapture, WasapiCapture};
|
||||
use crate::diarization::{Diarizer, SherpaDiarizer};
|
||||
use crate::error::WaResult;
|
||||
use crate::hardware::{HardwareDetector, WinHardwareDetector};
|
||||
use crate::models::*;
|
||||
use crate::notes::NotesRenderer;
|
||||
use crate::paths::{meeting_dir, settings_path, wa_root, whisper_model_file};
|
||||
use crate::paths::{
|
||||
diarization_embedding_model_file, diarization_segmentation_model_file, meeting_dir,
|
||||
settings_path, wa_root, whisper_model_file,
|
||||
};
|
||||
use crate::storage::{FinalizeMeeting, Meeting, NewMeeting};
|
||||
use crate::transcription::{
|
||||
models as model_catalog, run_streaming_worker, Transcriber, WhisperTranscriber,
|
||||
};
|
||||
use crate::{error::WaError, AppState, RecordingSession};
|
||||
use serde::Deserialize;
|
||||
use std::collections::HashMap;
|
||||
use std::path::PathBuf;
|
||||
use std::sync::{Arc, Mutex as StdMutex};
|
||||
use tauri::{AppHandle, Emitter, State};
|
||||
use tauri::{AppHandle, Emitter, Manager, State};
|
||||
|
||||
#[derive(Deserialize)]
|
||||
pub struct StartRecordingArgs {
|
||||
@@ -98,6 +103,54 @@ fn backend_for(settings: &Settings) -> BackendId {
|
||||
WinHardwareDetector.best(preferred).id
|
||||
}
|
||||
|
||||
/// Builds a `Diarizer` if both diarization models are installed — a fixed
|
||||
/// pair of well-known filenames, downloadable/removable via
|
||||
/// `diarization::models` and `list_diarization_models`/`download_model`/
|
||||
/// `remove_model` (T4.7). Returns `None` rather than erring — diarization is
|
||||
/// a provisional/refinement layer that a recording never depends on, same
|
||||
/// treatment as a missing hardware backend.
|
||||
fn diarizer_from_installed_models() -> Option<SherpaDiarizer> {
|
||||
let seg_model = diarization_segmentation_model_file();
|
||||
let emb_model = diarization_embedding_model_file();
|
||||
if !seg_model.exists() || !emb_model.exists() {
|
||||
return None;
|
||||
}
|
||||
match SherpaDiarizer::new(&seg_model, &emb_model) {
|
||||
Ok(d) => Some(d),
|
||||
Err(e) => {
|
||||
tracing::warn!("diarization models present but failed to load: {e}");
|
||||
None
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
/// The distinct speakers seen in `segments` so far, in first-appearance order,
|
||||
/// with any display names applied (T4.3/T4.4, FR-SPK-2/5). Falls back to the
|
||||
/// single pre-diarization "S1" placeholder if no segments exist yet.
|
||||
fn speaker_infos_from_segments(
|
||||
segments: &[TranscriptSegment],
|
||||
names: &HashMap<String, String>,
|
||||
) -> Vec<SpeakerInfo> {
|
||||
let mut seen = std::collections::HashSet::new();
|
||||
let mut out: Vec<SpeakerInfo> = segments
|
||||
.iter()
|
||||
.filter(|s| seen.insert(s.speaker.clone()))
|
||||
.map(|s| SpeakerInfo {
|
||||
label: s.speaker.clone(),
|
||||
display_name: names.get(&s.speaker).cloned(),
|
||||
participant_id: None,
|
||||
})
|
||||
.collect();
|
||||
if out.is_empty() {
|
||||
out.push(SpeakerInfo {
|
||||
label: "S1".to_string(),
|
||||
display_name: names.get("S1").cloned(),
|
||||
participant_id: None,
|
||||
});
|
||||
}
|
||||
out
|
||||
}
|
||||
|
||||
// ---- Recording lifecycle (Phase 1) ----
|
||||
|
||||
#[tauri::command]
|
||||
@@ -201,6 +254,79 @@ pub async fn start_recording(
|
||||
})
|
||||
.map_err(|e| WaError::new("transcription", e.to_string()))?;
|
||||
|
||||
// T4.3: cheap/provisional live diarization, skipped entirely (None) when
|
||||
// the diarization models aren't installed yet (T4.7) — never blocks
|
||||
// recording, same graceful-degradation treatment as a missing backend.
|
||||
// Loading the ONNX models is blocking I/O, so it runs off this async task.
|
||||
let diarizer: Option<Arc<dyn Diarizer>> =
|
||||
tauri::async_runtime::spawn_blocking(diarizer_from_installed_models)
|
||||
.await
|
||||
.ok()
|
||||
.flatten()
|
||||
.map(|d| Arc::new(d) as Arc<dyn Diarizer>);
|
||||
let speaker_names: Arc<StdMutex<HashMap<String, String>>> =
|
||||
Arc::new(StdMutex::new(HashMap::new()));
|
||||
|
||||
if let Some(diarizer) = diarizer.clone() {
|
||||
let app_for_diar = app.clone();
|
||||
let meeting_id_for_diar = meeting_id.clone();
|
||||
let wav_path_for_diar = wav_path.clone();
|
||||
let segments_for_diar = segments.clone();
|
||||
let names_for_diar = speaker_names.clone();
|
||||
tauri::async_runtime::spawn(async move {
|
||||
// ponytail: reprocesses the whole recording-so-far each tick
|
||||
// rather than incremental/windowed segmentation — sherpa-onnx's
|
||||
// offline Diarizer has no streaming primitive to build on, and
|
||||
// this is provisional preview only (the accurate pass runs once
|
||||
// at stop). Fine at meeting length and a 15s cadence; revisit
|
||||
// with real streaming segmentation if long meetings make it heavy.
|
||||
let mut ticker = tokio::time::interval(std::time::Duration::from_secs(15));
|
||||
ticker.tick().await; // interval's first tick fires immediately; skip it
|
||||
loop {
|
||||
ticker.tick().await;
|
||||
let still_active = app_for_diar
|
||||
.state::<AppState>()
|
||||
.session
|
||||
.lock()
|
||||
.await
|
||||
.as_ref()
|
||||
.is_some_and(|s| s.meeting_id == meeting_id_for_diar);
|
||||
if !still_active {
|
||||
break; // recording stopped (or a new one started) — nothing left to do
|
||||
}
|
||||
|
||||
let diarizer_for_pass = diarizer.clone();
|
||||
let wav_path = wav_path_for_diar.clone();
|
||||
let spans = tauri::async_runtime::spawn_blocking(move || {
|
||||
diarizer_for_pass.diarize(&wav_path)
|
||||
})
|
||||
.await;
|
||||
let spans = match spans {
|
||||
Ok(Ok(spans)) => spans,
|
||||
Ok(Err(e)) => {
|
||||
tracing::warn!("live diarization pass failed: {e}");
|
||||
continue;
|
||||
}
|
||||
Err(e) => {
|
||||
tracing::warn!("live diarization task failed: {e}");
|
||||
continue;
|
||||
}
|
||||
};
|
||||
|
||||
let speakers = {
|
||||
let mut segs = segments_for_diar.lock().unwrap_or_else(|e| e.into_inner());
|
||||
diarizer.assign(&mut segs, &spans);
|
||||
let names = names_for_diar.lock().unwrap_or_else(|e| e.into_inner());
|
||||
speaker_infos_from_segments(&segs, &names)
|
||||
};
|
||||
let _ = app_for_diar.emit(
|
||||
"diarization://updated",
|
||||
serde_json::json!({ "meetingId": meeting_id_for_diar, "speakers": speakers }),
|
||||
);
|
||||
}
|
||||
});
|
||||
}
|
||||
|
||||
*guard = Some(RecordingSession {
|
||||
meeting_id: meeting_id.clone(),
|
||||
capture,
|
||||
@@ -211,6 +337,8 @@ pub async fn start_recording(
|
||||
segments,
|
||||
active_backend,
|
||||
model_id,
|
||||
diarizer,
|
||||
speaker_names,
|
||||
});
|
||||
drop(guard);
|
||||
|
||||
@@ -255,18 +383,35 @@ pub async fn stop_recording(
|
||||
// guarantee it has fully drained the audio before we act on retention.
|
||||
let _ = session.transcription_worker.join();
|
||||
|
||||
let segments = session
|
||||
let mut segments = session
|
||||
.segments
|
||||
.lock()
|
||||
.map(|g| g.clone())
|
||||
.unwrap_or_default();
|
||||
let segment_count = segments.len();
|
||||
// No diarization until Phase 4 — every segment is provisionally "S1".
|
||||
let speakers = vec![SpeakerInfo {
|
||||
label: "S1".to_string(),
|
||||
display_name: None,
|
||||
participant_id: None,
|
||||
}];
|
||||
|
||||
// T4.1/4.2: one authoritative diarization pass over the now-complete
|
||||
// recording (ADR-0005's "post-stop pass"), refining whatever the live
|
||||
// provisional passes (T4.3) produced. Skipped if diarization models
|
||||
// aren't installed — `speaker_infos_from_segments` then falls back to
|
||||
// the single pre-diarization "S1" placeholder, same as before Phase 4.
|
||||
if let Some(diarizer) = session.diarizer.clone() {
|
||||
let diarizer_for_task = diarizer.clone();
|
||||
let wav_path = session.wav_path.clone();
|
||||
match tauri::async_runtime::spawn_blocking(move || diarizer_for_task.diarize(&wav_path))
|
||||
.await
|
||||
{
|
||||
Ok(Ok(spans)) => diarizer.assign(&mut segments, &spans),
|
||||
Ok(Err(e)) => tracing::warn!("final diarization pass failed: {e}"),
|
||||
Err(e) => tracing::warn!("final diarization task failed: {e}"),
|
||||
}
|
||||
}
|
||||
let speaker_names = session
|
||||
.speaker_names
|
||||
.lock()
|
||||
.map(|g| g.clone())
|
||||
.unwrap_or_default();
|
||||
let speakers = speaker_infos_from_segments(&segments, &speaker_names);
|
||||
let backend_used = session
|
||||
.active_backend
|
||||
.lock()
|
||||
@@ -395,6 +540,120 @@ pub async fn acknowledge_recording_consent() -> WaResult<()> {
|
||||
save_settings(&settings)
|
||||
}
|
||||
|
||||
// ---- Speakers (Phase 4) ----
|
||||
|
||||
/// Re-renders and persists `notes.md` from a finalized meeting's current
|
||||
/// (post-rename/post-merge) segments+speakers, and tells the frontend what
|
||||
/// changed (T4.5/4.6, FR-SPK-3/5). This is what keeps `export_meeting` — which
|
||||
/// just copies the already-rendered `notes.md` — in sync with naming changes
|
||||
/// made after the meeting ends; `transcript.json` itself is untouched.
|
||||
async fn refresh_notes_and_notify(
|
||||
app: &AppHandle,
|
||||
state: &State<'_, AppState>,
|
||||
meeting_id: &MeetingId,
|
||||
) -> WaResult<Vec<SpeakerInfo>> {
|
||||
let meeting = state
|
||||
.store
|
||||
.get_meeting(meeting_id)
|
||||
.await
|
||||
.map_err(|e| WaError::new("storage", e.to_string()))?;
|
||||
let notes_md =
|
||||
crate::notes::MarkdownNotes.to_markdown(&meeting.segments, &meeting.speakers, None);
|
||||
let _ = std::fs::write(meeting_dir(meeting_id).join("notes.md"), notes_md);
|
||||
let _ = app.emit(
|
||||
"diarization://updated",
|
||||
serde_json::json!({ "meetingId": meeting_id, "speakers": meeting.speakers }),
|
||||
);
|
||||
Ok(meeting.speakers)
|
||||
}
|
||||
|
||||
/// Name a speaker; applies to that speaker's past & future segments (T4.4,
|
||||
/// FR-SPK-2). Segments only ever carry the internal label ("S1"…) — never
|
||||
/// rewritten — so persisting the label→name mapping here is enough to cover
|
||||
/// both past and future segments once names are resolved at render time
|
||||
/// (FR-SPK-5). Works whether the meeting is still recording or already
|
||||
/// finalized.
|
||||
#[tauri::command]
|
||||
pub async fn rename_speaker(
|
||||
app: AppHandle,
|
||||
state: State<'_, AppState>,
|
||||
meeting_id: MeetingId,
|
||||
label: String,
|
||||
name: String,
|
||||
) -> WaResult<()> {
|
||||
state
|
||||
.store
|
||||
.rename_speaker(&meeting_id, &label, &name)
|
||||
.await
|
||||
.map_err(|e| WaError::new("storage", e.to_string()))?;
|
||||
|
||||
let guard = state.session.lock().await;
|
||||
match guard.as_ref().filter(|s| s.meeting_id == meeting_id) {
|
||||
Some(session) => {
|
||||
// Still recording: transcript.json/notes.md don't exist on disk
|
||||
// yet, so there's nothing to re-render — just refresh the live
|
||||
// in-memory view (finalize builds notes.md from this same map
|
||||
// at stop, T4.3/4.4).
|
||||
let names = {
|
||||
let mut names = session
|
||||
.speaker_names
|
||||
.lock()
|
||||
.unwrap_or_else(|e| e.into_inner());
|
||||
names.insert(label, name);
|
||||
names.clone()
|
||||
};
|
||||
let segments = session
|
||||
.segments
|
||||
.lock()
|
||||
.map(|g| g.clone())
|
||||
.unwrap_or_default();
|
||||
drop(guard);
|
||||
let speakers = speaker_infos_from_segments(&segments, &names);
|
||||
let _ = app.emit(
|
||||
"diarization://updated",
|
||||
serde_json::json!({ "meetingId": meeting_id, "speakers": speakers }),
|
||||
);
|
||||
}
|
||||
None => {
|
||||
drop(guard);
|
||||
refresh_notes_and_notify(&app, &state, &meeting_id).await?;
|
||||
}
|
||||
}
|
||||
Ok(())
|
||||
}
|
||||
|
||||
/// Fold over-split speakers into one canonical label (T4.5, FR-SPK-3) — e.g.
|
||||
/// diarization split one person into "S2" and "S3"; merging them shows one
|
||||
/// name and one grouped paragraph in notes/export. Post-meeting only: while
|
||||
/// still recording, the live provisional pass (T4.3) re-clusters from scratch
|
||||
/// every tick, so a label merged now could mean something else by the next
|
||||
/// tick.
|
||||
#[tauri::command]
|
||||
pub async fn merge_speakers(
|
||||
app: AppHandle,
|
||||
state: State<'_, AppState>,
|
||||
meeting_id: MeetingId,
|
||||
from: Vec<String>,
|
||||
into: String,
|
||||
) -> WaResult<()> {
|
||||
let guard = state.session.lock().await;
|
||||
if guard.as_ref().is_some_and(|s| s.meeting_id == meeting_id) {
|
||||
return Err(WaError::new(
|
||||
"recording",
|
||||
"cannot merge speakers while this meeting is still recording — wait until it's stopped",
|
||||
));
|
||||
}
|
||||
drop(guard);
|
||||
|
||||
state
|
||||
.store
|
||||
.merge_speakers(&meeting_id, &from, &into)
|
||||
.await
|
||||
.map_err(|e| WaError::new("storage", e.to_string()))?;
|
||||
refresh_notes_and_notify(&app, &state, &meeting_id).await?;
|
||||
Ok(())
|
||||
}
|
||||
|
||||
// ---- Hardware + models (Phase 3) ----
|
||||
|
||||
#[tauri::command]
|
||||
@@ -431,34 +690,55 @@ pub async fn list_models() -> WaResult<Vec<ModelInfo>> {
|
||||
Ok(model_catalog::list(&model_id_for(&settings)))
|
||||
}
|
||||
|
||||
/// The fixed segmentation+embedding pair (T4.7, FR-MODEL-1) — a separate
|
||||
/// command rather than folding into `list_models` because they're a fixed
|
||||
/// installable pair, not an interchangeable-size catalog like whisper's.
|
||||
#[tauri::command]
|
||||
pub async fn list_diarization_models() -> WaResult<Vec<ModelInfo>> {
|
||||
Ok(crate::diarization::models::list())
|
||||
}
|
||||
|
||||
#[derive(Deserialize)]
|
||||
pub struct DownloadModelArgs {
|
||||
pub kind: String, // "whisper" — diar-seg/diar-emb land in Phase 4 (T4.7)
|
||||
pub kind: String, // "whisper" | "diar-seg" | "diar-emb"
|
||||
pub id: String,
|
||||
}
|
||||
|
||||
#[tauri::command]
|
||||
pub async fn download_model(app: AppHandle, args: DownloadModelArgs) -> WaResult<()> {
|
||||
if args.kind != "whisper" {
|
||||
return Err(WaError::new(
|
||||
"model",
|
||||
format!("model kind '{}' is not available yet", args.kind),
|
||||
));
|
||||
}
|
||||
let id = args.id;
|
||||
let id_for_progress = id.clone();
|
||||
model_catalog::download(&id, move |received, total| {
|
||||
let on_progress = move |received: u64, total: Option<u64>| {
|
||||
let _ = app.emit(
|
||||
"model://progress",
|
||||
serde_json::json!({ "id": id_for_progress, "receivedBytes": received, "totalBytes": total }),
|
||||
);
|
||||
})
|
||||
.await
|
||||
.map_err(|e| WaError::new("model", e.to_string()))
|
||||
};
|
||||
match args.kind.as_str() {
|
||||
"whisper" => model_catalog::download(&id, on_progress)
|
||||
.await
|
||||
.map_err(|e| WaError::new("model", e.to_string())),
|
||||
"diar-seg" | "diar-emb" => crate::diarization::models::download(&id, on_progress)
|
||||
.await
|
||||
.map_err(|e| WaError::new("model", e.to_string())),
|
||||
other => Err(WaError::new(
|
||||
"model",
|
||||
format!("unknown model kind '{other}'"),
|
||||
)),
|
||||
}
|
||||
}
|
||||
|
||||
#[tauri::command]
|
||||
pub async fn remove_model(id: String) -> WaResult<()> {
|
||||
// No `kind` in this command's contract — disambiguate by catalog
|
||||
// membership instead; whisper/diarization ids never collide.
|
||||
if crate::diarization::models::list()
|
||||
.iter()
|
||||
.any(|m| m.id == id)
|
||||
{
|
||||
return crate::diarization::models::remove(&id)
|
||||
.map_err(|e| WaError::new("model", e.to_string()));
|
||||
}
|
||||
let settings = load_settings();
|
||||
model_catalog::remove(&id, &model_id_for(&settings))
|
||||
.map_err(|e| WaError::new("model", e.to_string()))
|
||||
@@ -913,3 +1193,54 @@ pub async fn update_settings(patch: serde_json::Value) -> WaResult<Settings> {
|
||||
pub async fn privacy_self_check() -> WaResult<serde_json::Value> {
|
||||
todo!("Phase 7 — privacy_self_check")
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::*;
|
||||
|
||||
fn segment(speaker: &str) -> TranscriptSegment {
|
||||
TranscriptSegment {
|
||||
id: 0,
|
||||
start_ms: 0,
|
||||
end_ms: 1000,
|
||||
speaker: speaker.to_string(),
|
||||
text: String::new(),
|
||||
confidence: None,
|
||||
interim: false,
|
||||
}
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn speaker_infos_lists_distinct_speakers_in_first_appearance_order() {
|
||||
let segments = vec![segment("S2"), segment("S1"), segment("S2")];
|
||||
let infos = speaker_infos_from_segments(&segments, &HashMap::new());
|
||||
let labels: Vec<&str> = infos.iter().map(|s| s.label.as_str()).collect();
|
||||
assert_eq!(labels, vec!["S2", "S1"]);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn speaker_infos_applies_display_names_by_label() {
|
||||
let segments = vec![segment("S1"), segment("S2")];
|
||||
let mut names = HashMap::new();
|
||||
names.insert("S1".to_string(), "Alice".to_string());
|
||||
let infos = speaker_infos_from_segments(&segments, &names);
|
||||
assert_eq!(infos[0].display_name.as_deref(), Some("Alice"));
|
||||
assert_eq!(infos[1].display_name, None); // S2 was never named
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn speaker_infos_falls_back_to_s1_placeholder_when_no_segments_yet() {
|
||||
let infos = speaker_infos_from_segments(&[], &HashMap::new());
|
||||
assert_eq!(infos.len(), 1);
|
||||
assert_eq!(infos[0].label, "S1");
|
||||
assert!(infos[0].display_name.is_none());
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn diarizer_is_none_when_models_are_not_installed() {
|
||||
// This test environment never has the fixed-path diarization models
|
||||
// installed (T4.7 will add real download/selection) — confirms the
|
||||
// graceful-degradation path a recording never blocks on (T4.3).
|
||||
assert!(diarizer_from_installed_models().is_none());
|
||||
}
|
||||
}
|
||||
|
||||
@@ -9,6 +9,8 @@
|
||||
use crate::models::{SpeakerSpan, TranscriptSegment};
|
||||
use std::path::Path;
|
||||
|
||||
pub mod models;
|
||||
|
||||
#[derive(Debug, thiserror::Error)]
|
||||
pub enum DiarError {
|
||||
#[error("model load failed: {0}")]
|
||||
@@ -24,17 +26,167 @@ pub trait Diarizer: Send + Sync {
|
||||
fn assign(&self, segments: &mut [TranscriptSegment], spans: &[SpeakerSpan]);
|
||||
}
|
||||
|
||||
/// Labels each segment with whichever span overlaps it most, in milliseconds
|
||||
/// (T4.2, FR-SPK-1). A segment with no overlapping span (e.g. it falls in a
|
||||
/// gap between spans) keeps its prior speaker label — the "S1" placeholder
|
||||
/// every segment starts with pre-diarization — rather than guessing.
|
||||
///
|
||||
/// Pure timestamp arithmetic, so it's engine-agnostic: every `Diarizer` impl
|
||||
/// can share it instead of reimplementing overlap math.
|
||||
pub fn assign_by_overlap(segments: &mut [TranscriptSegment], spans: &[SpeakerSpan]) {
|
||||
for segment in segments.iter_mut() {
|
||||
let best_span = spans
|
||||
.iter()
|
||||
.map(|span| {
|
||||
let overlap_start = segment.start_ms.max(span.start_ms);
|
||||
let overlap_end = segment.end_ms.min(span.end_ms);
|
||||
(overlap_end.saturating_sub(overlap_start), span)
|
||||
})
|
||||
.filter(|(overlap, _)| *overlap > 0)
|
||||
.max_by_key(|(overlap, _)| *overlap)
|
||||
.map(|(_, span)| span);
|
||||
|
||||
if let Some(span) = best_span {
|
||||
segment.speaker = span.speaker.clone();
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
/// sherpa-onnx-backed diarizer: pyannote segmentation + speaker-embedding +
|
||||
/// fast clustering (ADR-0005, T4.1). `Diarize::compute` needs `&mut self`; it's
|
||||
/// wrapped in a `Mutex` to satisfy `Diarizer: Sync` — diarization is a
|
||||
/// once-per-meeting post-pass (never a hot path), so lock contention is moot.
|
||||
#[cfg(feature = "diarization")]
|
||||
pub struct SherpaDiarizer;
|
||||
pub struct SherpaDiarizer {
|
||||
engine: std::sync::Mutex<sherpa_rs::diarize::Diarize>,
|
||||
}
|
||||
|
||||
#[cfg(feature = "diarization")]
|
||||
impl SherpaDiarizer {
|
||||
pub fn new(segmentation_model: &Path, embedding_model: &Path) -> Result<Self, DiarError> {
|
||||
let config = sherpa_rs::diarize::DiarizeConfig {
|
||||
// A meeting's speaker count isn't known ahead of time: <= 0 tells
|
||||
// sherpa-onnx to pick the cluster count itself from `threshold`
|
||||
// instead of forcing a fixed number of speakers.
|
||||
num_clusters: Some(-1),
|
||||
threshold: Some(0.5),
|
||||
..Default::default()
|
||||
};
|
||||
let engine = sherpa_rs::diarize::Diarize::new(segmentation_model, embedding_model, config)
|
||||
.map_err(|e| DiarError::Load(e.to_string()))?;
|
||||
Ok(Self {
|
||||
engine: std::sync::Mutex::new(engine),
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
#[cfg(feature = "diarization")]
|
||||
impl Diarizer for SherpaDiarizer {
|
||||
fn diarize(&self, _wav: &Path) -> Result<Vec<SpeakerSpan>, DiarError> {
|
||||
// T4.1: sherpa-onnx segmentation + embedding + clustering via FFI.
|
||||
todo!("Phase 4 — diarize")
|
||||
fn diarize(&self, wav: &Path) -> Result<Vec<SpeakerSpan>, DiarError> {
|
||||
let samples =
|
||||
crate::audio::read_wav_mono_16k(wav).map_err(|e| DiarError::Run(e.to_string()))?;
|
||||
let mut engine = self.engine.lock().unwrap_or_else(|e| e.into_inner());
|
||||
let segments = engine
|
||||
.compute(samples, None)
|
||||
.map_err(|e| DiarError::Run(e.to_string()))?;
|
||||
Ok(segments.into_iter().map(segment_to_span).collect())
|
||||
}
|
||||
fn assign(&self, _segments: &mut [TranscriptSegment], _spans: &[SpeakerSpan]) {
|
||||
// T4.2: timestamp-overlap alignment.
|
||||
todo!("Phase 4 — assign speakers to segments")
|
||||
|
||||
fn assign(&self, segments: &mut [TranscriptSegment], spans: &[SpeakerSpan]) {
|
||||
assign_by_overlap(segments, spans);
|
||||
}
|
||||
}
|
||||
|
||||
/// sherpa-onnx speaker indices are 0-based; WA's internal labels are 1-based ("S1"…).
|
||||
#[cfg(feature = "diarization")]
|
||||
fn segment_to_span(seg: sherpa_rs::diarize::Segment) -> SpeakerSpan {
|
||||
SpeakerSpan {
|
||||
start_ms: (seg.start.max(0.0) * 1000.0) as u64,
|
||||
end_ms: (seg.end.max(0.0) * 1000.0) as u64,
|
||||
speaker: format!("S{}", seg.speaker + 1),
|
||||
}
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod overlap_tests {
|
||||
use super::*;
|
||||
|
||||
fn segment(start_ms: u64, end_ms: u64) -> TranscriptSegment {
|
||||
TranscriptSegment {
|
||||
id: 0,
|
||||
start_ms,
|
||||
end_ms,
|
||||
speaker: "S1".to_string(),
|
||||
text: String::new(),
|
||||
confidence: None,
|
||||
interim: false,
|
||||
}
|
||||
}
|
||||
|
||||
fn span(start_ms: u64, end_ms: u64, speaker: &str) -> SpeakerSpan {
|
||||
SpeakerSpan {
|
||||
start_ms,
|
||||
end_ms,
|
||||
speaker: speaker.to_string(),
|
||||
}
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn assigns_the_span_that_overlaps_most() {
|
||||
let mut segments = vec![segment(0, 1000), segment(1000, 2000)];
|
||||
let spans = vec![span(0, 1000, "S1"), span(1000, 2000, "S2")];
|
||||
assign_by_overlap(&mut segments, &spans);
|
||||
assert_eq!(segments[0].speaker, "S1");
|
||||
assert_eq!(segments[1].speaker, "S2");
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn a_segment_straddling_two_spans_picks_the_larger_overlap() {
|
||||
// 700ms in S1's span (300-1000), 300ms in S2's span (1000-1300).
|
||||
let mut segments = vec![segment(300, 1300)];
|
||||
let spans = vec![span(0, 1000, "S1"), span(1000, 2000, "S2")];
|
||||
assign_by_overlap(&mut segments, &spans);
|
||||
assert_eq!(segments[0].speaker, "S1");
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn a_segment_with_no_overlapping_span_keeps_its_prior_label() {
|
||||
let mut segments = vec![segment(5000, 6000)];
|
||||
segments[0].speaker = "S9".to_string(); // distinct from any span below
|
||||
let spans = vec![span(0, 1000, "S1")];
|
||||
assign_by_overlap(&mut segments, &spans);
|
||||
assert_eq!(segments[0].speaker, "S9"); // untouched, not overwritten with a guess
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn no_spans_at_all_leaves_segments_untouched() {
|
||||
let mut segments = vec![segment(0, 1000)];
|
||||
assign_by_overlap(&mut segments, &[]);
|
||||
assert_eq!(segments[0].speaker, "S1");
|
||||
}
|
||||
}
|
||||
|
||||
#[cfg(all(test, feature = "diarization"))]
|
||||
mod tests {
|
||||
use super::*;
|
||||
|
||||
#[test]
|
||||
fn segment_to_span_maps_0_based_speaker_index_to_1_based_label() {
|
||||
let seg = sherpa_rs::diarize::Segment {
|
||||
start: 1.5,
|
||||
end: 3.25,
|
||||
speaker: 0,
|
||||
};
|
||||
let span = segment_to_span(seg);
|
||||
assert_eq!(span.start_ms, 1500);
|
||||
assert_eq!(span.end_ms, 3250);
|
||||
assert_eq!(span.speaker, "S1");
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn new_surfaces_a_load_error_for_missing_models_instead_of_panicking() {
|
||||
let result =
|
||||
SherpaDiarizer::new(Path::new("no-such-seg.onnx"), Path::new("no-such-emb.onnx"));
|
||||
assert!(matches!(result, Err(DiarError::Load(_))));
|
||||
}
|
||||
}
|
||||
|
||||
@@ -0,0 +1,134 @@
|
||||
//! Diarization model catalog + install/remove (Phase 4, T4.7, FR-MODEL-1).
|
||||
//!
|
||||
//! ponytail: two fixed models (segmentation + embedding, ADR-0005), not a
|
||||
//! fetched index — same lazy-correct call as the whisper catalog (T3.7).
|
||||
//! Mirrors the pyannote-segmentation-3.0 + 3D-Speaker ERes2Net (English)
|
||||
//! models sherpa-onnx's own docs use for offline speaker diarization.
|
||||
|
||||
use crate::models::ModelInfo;
|
||||
use crate::paths::{
|
||||
diarization_embedding_model_file, diarization_segmentation_model_file, models_dir,
|
||||
};
|
||||
use futures_util::StreamExt;
|
||||
use std::path::PathBuf;
|
||||
|
||||
struct Catalog {
|
||||
id: &'static str,
|
||||
label: &'static str,
|
||||
size_mb: u32,
|
||||
url: &'static str,
|
||||
dest: fn() -> PathBuf,
|
||||
}
|
||||
|
||||
const CATALOG: &[Catalog] = &[
|
||||
Catalog {
|
||||
id: "seg-pyannote-3.0",
|
||||
label: "Speaker segmentation (pyannote 3.0)",
|
||||
size_mb: 6,
|
||||
url: "https://huggingface.co/csukuangfj/sherpa-onnx-pyannote-segmentation-3-0/resolve/main/model.onnx",
|
||||
dest: diarization_segmentation_model_file,
|
||||
},
|
||||
Catalog {
|
||||
id: "spk-eres2net",
|
||||
label: "Speaker embedding (3D-Speaker ERes2Net, English)",
|
||||
size_mb: 27,
|
||||
url: "https://huggingface.co/csukuangfj/speaker-embedding-models/resolve/main/3dspeaker_speech_eres2net_sv_en_voxceleb_16k.onnx",
|
||||
dest: diarization_embedding_model_file,
|
||||
},
|
||||
];
|
||||
|
||||
fn find(id: &str) -> Option<&'static Catalog> {
|
||||
CATALOG.iter().find(|m| m.id == id)
|
||||
}
|
||||
|
||||
pub fn list() -> Vec<ModelInfo> {
|
||||
CATALOG
|
||||
.iter()
|
||||
.map(|m| ModelInfo {
|
||||
id: m.id.to_string(),
|
||||
label: m.label.to_string(),
|
||||
size_mb: m.size_mb,
|
||||
installed: (m.dest)().exists(),
|
||||
// Both models are always "active" once installed — diarization
|
||||
// has no interchangeable-size picker like whisper's (yet).
|
||||
active: true,
|
||||
})
|
||||
.collect()
|
||||
}
|
||||
|
||||
#[derive(Debug, thiserror::Error)]
|
||||
pub enum ModelError {
|
||||
#[error("unknown diarization model id: {0}")]
|
||||
UnknownId(String),
|
||||
#[error("download failed: {0}")]
|
||||
Download(String),
|
||||
#[error("io error: {0}")]
|
||||
Io(#[from] std::io::Error),
|
||||
}
|
||||
|
||||
/// Downloads a diarization model, reporting progress via `on_progress`.
|
||||
/// Writes to a `.part` file first so a crash/cancel never leaves a truncated
|
||||
/// model that `diarizer_from_installed_models` would treat as installed.
|
||||
pub async fn download(
|
||||
id: &str,
|
||||
mut on_progress: impl FnMut(u64, Option<u64>),
|
||||
) -> Result<(), ModelError> {
|
||||
let catalog = find(id).ok_or_else(|| ModelError::UnknownId(id.to_string()))?;
|
||||
std::fs::create_dir_all(models_dir())?;
|
||||
let dest = (catalog.dest)();
|
||||
let tmp = dest.with_extension("part");
|
||||
|
||||
let resp = reqwest::get(catalog.url)
|
||||
.await
|
||||
.map_err(|e| ModelError::Download(e.to_string()))?;
|
||||
if !resp.status().is_success() {
|
||||
return Err(ModelError::Download(format!("HTTP {}", resp.status())));
|
||||
}
|
||||
let total = resp.content_length();
|
||||
let mut received: u64 = 0;
|
||||
let mut file = std::fs::File::create(&tmp)?;
|
||||
let mut stream = resp.bytes_stream();
|
||||
while let Some(chunk) = stream.next().await {
|
||||
let chunk = chunk.map_err(|e| ModelError::Download(e.to_string()))?;
|
||||
std::io::Write::write_all(&mut file, &chunk)?;
|
||||
received += chunk.len() as u64;
|
||||
on_progress(received, total);
|
||||
}
|
||||
drop(file);
|
||||
std::fs::rename(&tmp, &dest)?;
|
||||
Ok(())
|
||||
}
|
||||
|
||||
pub fn remove(id: &str) -> Result<(), ModelError> {
|
||||
let catalog = find(id).ok_or_else(|| ModelError::UnknownId(id.to_string()))?;
|
||||
let path = (catalog.dest)();
|
||||
if path.exists() {
|
||||
std::fs::remove_file(path)?;
|
||||
}
|
||||
Ok(())
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::*;
|
||||
|
||||
#[test]
|
||||
fn list_returns_exactly_the_segmentation_and_embedding_pair() {
|
||||
let models = list();
|
||||
assert_eq!(models.len(), 2);
|
||||
assert!(models.iter().any(|m| m.id == "seg-pyannote-3.0"));
|
||||
assert!(models.iter().any(|m| m.id == "spk-eres2net"));
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn download_rejects_an_unknown_id_without_touching_disk() {
|
||||
let err = download("no-such-model", |_, _| {}).await.unwrap_err();
|
||||
assert!(matches!(err, ModelError::UnknownId(_)));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn remove_rejects_an_unknown_id() {
|
||||
let err = remove("no-such-model").unwrap_err();
|
||||
assert!(matches!(err, ModelError::UnknownId(_)));
|
||||
}
|
||||
}
|
||||
@@ -56,6 +56,14 @@ pub struct RecordingSession {
|
||||
/// can pick a different model than `Settings.whisper_model`).
|
||||
pub active_backend: Arc<StdMutex<models::BackendId>>,
|
||||
pub model_id: String,
|
||||
/// `None` when diarization models aren't installed yet (T4.7) — live
|
||||
/// provisional turns and the final post-stop pass are both skipped, same
|
||||
/// graceful-degradation treatment as a missing hardware backend (T4.3).
|
||||
pub diarizer: Option<Arc<dyn diarization::Diarizer>>,
|
||||
/// label ("S1"…) -> user-given display name, settable mid-recording
|
||||
/// (T4.4, FR-SPK-2). Never rewritten onto segments (FR-SPK-5); resolved
|
||||
/// at render/finalize time instead.
|
||||
pub speaker_names: Arc<StdMutex<std::collections::HashMap<String, String>>>,
|
||||
}
|
||||
|
||||
/// Wraps the tray icon so it can be looked up from commands to update its
|
||||
@@ -120,6 +128,7 @@ pub fn run() {
|
||||
commands::hardware_status,
|
||||
commands::set_preferred_backend,
|
||||
commands::list_models,
|
||||
commands::list_diarization_models,
|
||||
commands::download_model,
|
||||
commands::remove_model,
|
||||
commands::reprocess_transcript,
|
||||
@@ -128,6 +137,8 @@ pub fn run() {
|
||||
commands::delete_meeting,
|
||||
commands::update_notes,
|
||||
commands::export_meeting,
|
||||
commands::rename_speaker,
|
||||
commands::merge_speakers,
|
||||
commands::llm_status,
|
||||
commands::set_llm_provider,
|
||||
commands::generate_summary,
|
||||
|
||||
@@ -43,3 +43,12 @@ pub fn whisper_model_file(id: &str) -> PathBuf {
|
||||
pub fn whisper_model_path() -> PathBuf {
|
||||
whisper_model_file(DEFAULT_WHISPER_MODEL)
|
||||
}
|
||||
|
||||
/// Fixed filenames pending T4.7 (diarization model management/selection in Settings).
|
||||
pub fn diarization_segmentation_model_file() -> PathBuf {
|
||||
models_dir().join("seg-pyannote-3.0.onnx")
|
||||
}
|
||||
|
||||
pub fn diarization_embedding_model_file() -> PathBuf {
|
||||
models_dir().join("spk-eres2net.onnx")
|
||||
}
|
||||
|
||||
+124
-13
@@ -10,6 +10,7 @@ use async_trait::async_trait;
|
||||
use serde::{Deserialize, Serialize};
|
||||
use sqlx::sqlite::{SqliteConnectOptions, SqlitePool, SqlitePoolOptions};
|
||||
use sqlx::Row;
|
||||
use std::collections::HashMap;
|
||||
use std::path::PathBuf;
|
||||
use std::time::{SystemTime, UNIX_EPOCH};
|
||||
|
||||
@@ -82,6 +83,24 @@ pub trait Store: Send + Sync {
|
||||
async fn get_meeting(&self, id: &MeetingId) -> Result<Meeting, StoreError>;
|
||||
async fn delete_meeting(&self, id: &MeetingId) -> Result<(), StoreError>;
|
||||
async fn update_notes(&self, id: &MeetingId, markdown: &str) -> Result<(), StoreError>;
|
||||
/// Set (or create) a speaker's display name; works whether or not the
|
||||
/// meeting has finalized yet (T4.4, FR-SPK-2/5).
|
||||
async fn rename_speaker(
|
||||
&self,
|
||||
id: &MeetingId,
|
||||
label: &str,
|
||||
name: &str,
|
||||
) -> Result<(), StoreError>;
|
||||
/// Fold over-split speaker labels into one canonical label (T4.5,
|
||||
/// FR-SPK-3). Segment speaker IDs in storage are never rewritten
|
||||
/// (FR-SPK-5) — `get_meeting` resolves `from` labels to `into` when it
|
||||
/// reads segments/speakers back.
|
||||
async fn merge_speakers(
|
||||
&self,
|
||||
id: &MeetingId,
|
||||
from: &[String],
|
||||
into: &str,
|
||||
) -> Result<(), StoreError>;
|
||||
/// Full-text search across transcripts + notes (Phase 8, FR-SEARCH-1).
|
||||
async fn search(&self, query: &str) -> Result<Vec<MeetingListItem>, StoreError>;
|
||||
/// Startup reconcile: meetings with audio but no finalized transcript (FR-REL-1).
|
||||
@@ -116,6 +135,42 @@ impl SqliteStore {
|
||||
}
|
||||
}
|
||||
|
||||
impl SqliteStore {
|
||||
async fn upsert_speaker(
|
||||
&self,
|
||||
meeting_id: &MeetingId,
|
||||
label: &str,
|
||||
display_name: Option<&str>,
|
||||
) -> Result<(), StoreError> {
|
||||
sqlx::query(
|
||||
"INSERT INTO speakers (id, meeting_id, label, display_name)
|
||||
VALUES (?, ?, ?, ?)
|
||||
ON CONFLICT(meeting_id, label) DO UPDATE SET display_name = excluded.display_name",
|
||||
)
|
||||
.bind(uuid::Uuid::new_v4().to_string())
|
||||
.bind(meeting_id)
|
||||
.bind(label)
|
||||
.bind(display_name)
|
||||
.execute(&self.pool)
|
||||
.await?;
|
||||
Ok(())
|
||||
}
|
||||
}
|
||||
|
||||
/// Follows a `label -> merged_into` chain to its canonical label. Capped at 8
|
||||
/// hops so a stale/cyclic mapping (shouldn't happen, but merges are
|
||||
/// user-driven data) can't loop forever; merges normally resolve in one hop.
|
||||
fn resolve_canonical<'a>(label: &'a str, merge_map: &'a HashMap<String, String>) -> &'a str {
|
||||
let mut current = label;
|
||||
for _ in 0..8 {
|
||||
match merge_map.get(current) {
|
||||
Some(next) if next != current => current = next.as_str(),
|
||||
_ => break,
|
||||
}
|
||||
}
|
||||
current
|
||||
}
|
||||
|
||||
fn now_unix() -> i64 {
|
||||
SystemTime::now()
|
||||
.duration_since(UNIX_EPOCH)
|
||||
@@ -194,17 +249,8 @@ impl Store for SqliteStore {
|
||||
.await?;
|
||||
|
||||
for speaker in &s.speakers {
|
||||
sqlx::query(
|
||||
"INSERT INTO speakers (id, meeting_id, label, display_name)
|
||||
VALUES (?, ?, ?, ?)
|
||||
ON CONFLICT(meeting_id, label) DO UPDATE SET display_name = excluded.display_name",
|
||||
)
|
||||
.bind(uuid::Uuid::new_v4().to_string())
|
||||
.bind(id)
|
||||
.bind(&speaker.label)
|
||||
.bind(&speaker.display_name)
|
||||
.execute(&self.pool)
|
||||
.await?;
|
||||
self.upsert_speaker(id, &speaker.label, speaker.display_name.as_deref())
|
||||
.await?;
|
||||
}
|
||||
|
||||
let transcript = TranscriptFile {
|
||||
@@ -261,13 +307,24 @@ impl Store for SqliteStore {
|
||||
.ok_or_else(|| StoreError::NotFound(id.clone()))?;
|
||||
|
||||
let speaker_rows = sqlx::query(
|
||||
"SELECT label, display_name FROM speakers WHERE meeting_id = ? ORDER BY label",
|
||||
"SELECT label, display_name, merged_into FROM speakers WHERE meeting_id = ? ORDER BY label",
|
||||
)
|
||||
.bind(id)
|
||||
.fetch_all(&self.pool)
|
||||
.await?;
|
||||
// T4.5: rows with merged_into set are folded away — only canonical
|
||||
// speakers are returned, but their raw label still resolves through
|
||||
// `merge_map` below so folded segments render under the right one.
|
||||
let merge_map: HashMap<String, String> = speaker_rows
|
||||
.iter()
|
||||
.filter_map(|r| {
|
||||
let merged_into: Option<String> = r.get("merged_into");
|
||||
merged_into.map(|into| (r.get::<String, _>("label"), into))
|
||||
})
|
||||
.collect();
|
||||
let speakers: Vec<SpeakerInfo> = speaker_rows
|
||||
.iter()
|
||||
.filter(|r| r.get::<Option<String>, _>("merged_into").is_none())
|
||||
.map(|r| SpeakerInfo {
|
||||
label: r.get("label"),
|
||||
display_name: r.get("display_name"),
|
||||
@@ -276,11 +333,19 @@ impl Store for SqliteStore {
|
||||
.collect();
|
||||
|
||||
let folder = paths::meeting_dir(id);
|
||||
let segments = std::fs::read_to_string(folder.join("transcript.json"))
|
||||
let mut segments = std::fs::read_to_string(folder.join("transcript.json"))
|
||||
.ok()
|
||||
.and_then(|s| serde_json::from_str::<TranscriptFile>(&s).ok())
|
||||
.map(|t| t.segments)
|
||||
.unwrap_or_default();
|
||||
if !merge_map.is_empty() {
|
||||
// Resolution only touches this in-memory copy — transcript.json
|
||||
// on disk keeps its raw labels, regenerable and non-destructive
|
||||
// (FR-SPK-5), same as display names.
|
||||
for seg in &mut segments {
|
||||
seg.speaker = resolve_canonical(&seg.speaker, &merge_map).to_string();
|
||||
}
|
||||
}
|
||||
let notes_markdown = std::fs::read_to_string(folder.join("notes.md")).unwrap_or_default();
|
||||
|
||||
Ok(Meeting {
|
||||
@@ -322,6 +387,52 @@ impl Store for SqliteStore {
|
||||
Ok(())
|
||||
}
|
||||
|
||||
async fn rename_speaker(
|
||||
&self,
|
||||
id: &MeetingId,
|
||||
label: &str,
|
||||
name: &str,
|
||||
) -> Result<(), StoreError> {
|
||||
self.upsert_speaker(id, label, Some(name)).await
|
||||
}
|
||||
|
||||
async fn merge_speakers(
|
||||
&self,
|
||||
id: &MeetingId,
|
||||
from: &[String],
|
||||
into: &str,
|
||||
) -> Result<(), StoreError> {
|
||||
// Ensure the canonical label has a row, without clobbering a name it
|
||||
// may already have (plain upsert_speaker would overwrite display_name
|
||||
// with `None` if `into` hasn't been named yet).
|
||||
sqlx::query(
|
||||
"INSERT INTO speakers (id, meeting_id, label) VALUES (?, ?, ?)
|
||||
ON CONFLICT(meeting_id, label) DO NOTHING",
|
||||
)
|
||||
.bind(uuid::Uuid::new_v4().to_string())
|
||||
.bind(id)
|
||||
.bind(into)
|
||||
.execute(&self.pool)
|
||||
.await?;
|
||||
|
||||
for label in from {
|
||||
if label == into {
|
||||
continue; // merging a label into itself is a no-op
|
||||
}
|
||||
sqlx::query(
|
||||
"INSERT INTO speakers (id, meeting_id, label, merged_into) VALUES (?, ?, ?, ?)
|
||||
ON CONFLICT(meeting_id, label) DO UPDATE SET merged_into = excluded.merged_into",
|
||||
)
|
||||
.bind(uuid::Uuid::new_v4().to_string())
|
||||
.bind(id)
|
||||
.bind(label)
|
||||
.bind(into)
|
||||
.execute(&self.pool)
|
||||
.await?;
|
||||
}
|
||||
Ok(())
|
||||
}
|
||||
|
||||
async fn search(&self, _query: &str) -> Result<Vec<MeetingListItem>, StoreError> {
|
||||
todo!("Phase 8 — FTS5 search")
|
||||
}
|
||||
|
||||
Reference in New Issue
Block a user