f5b99106b1
This commit introduces the capability to configure user-specific Whisper settings, including language, hotwords, and initial prompts. These settings are stored in `users.toml` and are passed to the Whisper service for more tailored transcription results. New unit tests have been added to verify that the `transcribe` function correctly forwards these optional settings to the Whisper client and omits them when they are not provided. Additionally, a new directory `tests/fixtures/dictations` has been created to store M4A audio files, their expected transcripts, and associated hotword files. This serves as a regression corpus for the Whisper pipeline. A README file explains the structure and conventions for adding new fixtures. Several new fixtures have been added, covering various medical domains and transcription scenarios.
351 lines
11 KiB
Rust
351 lines
11 KiB
Rust
use std::path::PathBuf;
|
|
use std::time::Duration;
|
|
|
|
use doctate_server::transcribe;
|
|
use doctate_server::transcribe::ffmpeg::remux_faststart;
|
|
use doctate_server::transcribe::ollama::{generate_oneliner, OllamaError};
|
|
use doctate_server::transcribe::recovery::scan_and_enqueue;
|
|
use doctate_server::config::WhisperUserSettings;
|
|
use doctate_server::transcribe::whisper::{transcribe, WhisperError};
|
|
use serde_json::json;
|
|
use wiremock::matchers::{body_partial_json, method, path, query_param};
|
|
use wiremock::{Mock, MockServer, ResponseTemplate};
|
|
|
|
fn fixture(name: &str) -> PathBuf {
|
|
PathBuf::from(env!("CARGO_MANIFEST_DIR"))
|
|
.join("tests/fixtures")
|
|
.join(name)
|
|
}
|
|
|
|
#[tokio::test]
|
|
async fn ffmpeg_remux_produces_valid_output() {
|
|
let input = fixture("sample.m4a");
|
|
assert!(input.exists(), "fixture missing: {}", input.display());
|
|
|
|
let output = remux_faststart(&input)
|
|
.await
|
|
.expect("remux failed");
|
|
|
|
let meta = std::fs::metadata(output.path()).expect("output missing");
|
|
assert!(meta.len() > 0, "output file is empty");
|
|
|
|
// Sanity check: ffprobe should still recognize it as an m4a/mp4 container.
|
|
let probe = std::process::Command::new("ffprobe")
|
|
.arg("-hide_banner")
|
|
.arg("-loglevel")
|
|
.arg("error")
|
|
.arg("-show_entries")
|
|
.arg("format=format_name")
|
|
.arg("-of")
|
|
.arg("default=nw=1:nk=1")
|
|
.arg(output.path())
|
|
.output()
|
|
.expect("ffprobe spawn failed");
|
|
assert!(probe.status.success(), "ffprobe failed");
|
|
let fmt = String::from_utf8_lossy(&probe.stdout);
|
|
assert!(fmt.contains("mp4") || fmt.contains("m4a"), "unexpected format: {fmt}");
|
|
}
|
|
|
|
#[tokio::test]
|
|
async fn ffmpeg_remux_fails_on_missing_input() {
|
|
let input = PathBuf::from("/nonexistent/does-not-exist.m4a");
|
|
let result = remux_faststart(&input).await;
|
|
assert!(result.is_err(), "expected error for missing input");
|
|
}
|
|
|
|
#[tokio::test]
|
|
async fn whisper_client_posts_multipart_and_returns_text() {
|
|
let server = MockServer::start().await;
|
|
let expected = "Hallo Welt, das ist ein Test.";
|
|
|
|
Mock::given(method("POST"))
|
|
.and(path("/asr"))
|
|
.and(query_param("output", "txt"))
|
|
.and(query_param("language", "de"))
|
|
.respond_with(ResponseTemplate::new(200).set_body_string(expected))
|
|
.expect(1)
|
|
.mount(&server)
|
|
.await;
|
|
|
|
let client = reqwest::Client::new();
|
|
let text = transcribe(
|
|
&client,
|
|
&server.uri(),
|
|
&fixture("sample.m4a"),
|
|
Duration::from_secs(10),
|
|
&WhisperUserSettings::default(),
|
|
)
|
|
.await
|
|
.expect("transcribe failed");
|
|
|
|
assert_eq!(text, expected);
|
|
}
|
|
|
|
#[tokio::test]
|
|
async fn whisper_client_returns_status_error_on_500() {
|
|
let server = MockServer::start().await;
|
|
|
|
Mock::given(method("POST"))
|
|
.and(path("/asr"))
|
|
.respond_with(ResponseTemplate::new(500).set_body_string("boom"))
|
|
.mount(&server)
|
|
.await;
|
|
|
|
let client = reqwest::Client::new();
|
|
let err = transcribe(
|
|
&client,
|
|
&server.uri(),
|
|
&fixture("sample.m4a"),
|
|
Duration::from_secs(10),
|
|
&WhisperUserSettings::default(),
|
|
)
|
|
.await
|
|
.expect_err("expected error");
|
|
|
|
match err {
|
|
WhisperError::Status { status, body } => {
|
|
assert_eq!(status, 500);
|
|
assert_eq!(body, "boom");
|
|
}
|
|
other => panic!("unexpected error variant: {other:?}"),
|
|
}
|
|
}
|
|
|
|
#[tokio::test]
|
|
async fn whisper_client_times_out() {
|
|
let server = MockServer::start().await;
|
|
|
|
Mock::given(method("POST"))
|
|
.and(path("/asr"))
|
|
.respond_with(ResponseTemplate::new(200).set_delay(Duration::from_secs(2)))
|
|
.mount(&server)
|
|
.await;
|
|
|
|
let client = reqwest::Client::new();
|
|
let err = transcribe(
|
|
&client,
|
|
&server.uri(),
|
|
&fixture("sample.m4a"),
|
|
Duration::from_millis(100),
|
|
&WhisperUserSettings::default(),
|
|
)
|
|
.await
|
|
.expect_err("expected timeout");
|
|
|
|
assert!(matches!(err, WhisperError::Http(_)), "expected Http error, got {err:?}");
|
|
}
|
|
|
|
#[tokio::test]
|
|
async fn whisper_client_forwards_hotwords_and_prompt_when_set() {
|
|
let server = MockServer::start().await;
|
|
|
|
// Accept any POST; inspect captured body below. language query param is
|
|
// still matched strictly to verify override from settings.
|
|
Mock::given(method("POST"))
|
|
.and(path("/asr"))
|
|
.and(query_param("language", "en"))
|
|
.respond_with(ResponseTemplate::new(200).set_body_string("ok"))
|
|
.expect(1)
|
|
.mount(&server)
|
|
.await;
|
|
|
|
let client = reqwest::Client::new();
|
|
let settings = WhisperUserSettings {
|
|
language: Some("en".into()),
|
|
hotwords: Some("HOCM Valsalva".into()),
|
|
initial_prompt: Some("Kardiologie".into()),
|
|
};
|
|
transcribe(
|
|
&client,
|
|
&server.uri(),
|
|
&fixture("sample.m4a"),
|
|
Duration::from_secs(10),
|
|
&settings,
|
|
)
|
|
.await
|
|
.expect("transcribe failed");
|
|
|
|
let received = server.received_requests().await.unwrap();
|
|
let body = String::from_utf8_lossy(&received[0].body);
|
|
assert!(body.contains("name=\"hotwords\""), "hotwords part missing");
|
|
assert!(body.contains("HOCM Valsalva"), "hotwords value missing: {body}");
|
|
assert!(body.contains("name=\"initial_prompt\""), "initial_prompt part missing");
|
|
assert!(body.contains("Kardiologie"), "initial_prompt value missing");
|
|
}
|
|
|
|
#[tokio::test]
|
|
async fn whisper_client_omits_empty_optional_fields() {
|
|
let server = MockServer::start().await;
|
|
|
|
// Default settings = all None → only `audio_file` part, language defaults
|
|
// to `de`. Assert that neither optional field appears in the body.
|
|
Mock::given(method("POST"))
|
|
.and(path("/asr"))
|
|
.and(query_param("language", "de"))
|
|
.respond_with(ResponseTemplate::new(200).set_body_string("ok"))
|
|
.expect(1)
|
|
.mount(&server)
|
|
.await;
|
|
|
|
let client = reqwest::Client::new();
|
|
transcribe(
|
|
&client,
|
|
&server.uri(),
|
|
&fixture("sample.m4a"),
|
|
Duration::from_secs(10),
|
|
&WhisperUserSettings::default(),
|
|
)
|
|
.await
|
|
.expect("transcribe failed");
|
|
|
|
// Inspect the captured request to confirm absence of the optional parts.
|
|
let received = server.received_requests().await.unwrap();
|
|
let body = String::from_utf8_lossy(&received[0].body);
|
|
assert!(!body.contains("name=\"hotwords\""), "hotwords part leaked: {body}");
|
|
assert!(!body.contains("name=\"initial_prompt\""), "initial_prompt part leaked: {body}");
|
|
}
|
|
|
|
#[tokio::test]
|
|
async fn recovery_enqueues_only_pending_recordings() {
|
|
let tmp = tempfile::tempdir().unwrap();
|
|
let root = tmp.path();
|
|
|
|
// User 1, open case with two .m4a — one pending, one already transcribed.
|
|
let case_a = root.join("dr_a/open/aaaa");
|
|
std::fs::create_dir_all(&case_a).unwrap();
|
|
std::fs::write(case_a.join("2026-04-10T10-00-00Z.m4a"), b"x").unwrap();
|
|
std::fs::write(case_a.join("2026-04-11T10-00-00Z.m4a"), b"x").unwrap();
|
|
std::fs::write(case_a.join("2026-04-11T10-00-00Z.transcript.txt"), b"done").unwrap();
|
|
|
|
// User 2, open case with one pending .m4a.
|
|
let case_b = root.join("dr_b/open/bbbb");
|
|
std::fs::create_dir_all(&case_b).unwrap();
|
|
std::fs::write(case_b.join("2026-04-12T10-00-00Z.m4a"), b"x").unwrap();
|
|
|
|
// done/ cases must be ignored (not re-transcribed).
|
|
let case_c = root.join("dr_a/done/cccc");
|
|
std::fs::create_dir_all(&case_c).unwrap();
|
|
std::fs::write(case_c.join("2026-04-09T10-00-00Z.m4a"), b"x").unwrap();
|
|
|
|
let (tx, mut rx) = transcribe::channel();
|
|
scan_and_enqueue(root, &tx).await;
|
|
drop(tx); // close channel so the loop below terminates
|
|
|
|
let mut received = Vec::new();
|
|
while let Some(job) = rx.recv().await {
|
|
received.push(job.audio_path);
|
|
}
|
|
|
|
assert_eq!(received.len(), 2, "got: {received:?}");
|
|
// Oldest first:
|
|
assert!(received[0].ends_with("2026-04-10T10-00-00Z.m4a"));
|
|
assert!(received[1].ends_with("2026-04-12T10-00-00Z.m4a"));
|
|
}
|
|
|
|
// -------------------- Ollama client --------------------
|
|
|
|
const MODEL: &str = "gemma3:4b";
|
|
|
|
fn ok_chat_body(content: &str) -> serde_json::Value {
|
|
json!({
|
|
"model": MODEL,
|
|
"message": {"role": "assistant", "content": content},
|
|
"done": true
|
|
})
|
|
}
|
|
|
|
#[tokio::test]
|
|
async fn ollama_oneliner_parses_chat_response() {
|
|
let server = MockServer::start().await;
|
|
|
|
Mock::given(method("POST"))
|
|
.and(path("/api/chat"))
|
|
.respond_with(ResponseTemplate::new(200).set_body_json(ok_chat_body("HOCM mit Septum-Hypertrophie")))
|
|
.expect(1)
|
|
.mount(&server)
|
|
.await;
|
|
|
|
let client = reqwest::Client::new();
|
|
let out = generate_oneliner(&client, &server.uri(), MODEL, 0, "irrelevant", Duration::from_secs(5))
|
|
.await
|
|
.expect("call failed");
|
|
|
|
assert_eq!(out, "HOCM mit Septum-Hypertrophie");
|
|
}
|
|
|
|
#[tokio::test]
|
|
async fn ollama_oneliner_normalizes_response() {
|
|
let server = MockServer::start().await;
|
|
Mock::given(method("POST"))
|
|
.and(path("/api/chat"))
|
|
.respond_with(ResponseTemplate::new(200).set_body_json(ok_chat_body("\"Kniegelenk re., V.a. Meniskus.\"\nSonstiges")))
|
|
.mount(&server)
|
|
.await;
|
|
|
|
let client = reqwest::Client::new();
|
|
let out = generate_oneliner(&client, &server.uri(), MODEL, 0, "x", Duration::from_secs(5))
|
|
.await
|
|
.unwrap();
|
|
assert_eq!(out, "Kniegelenk re., V.a. Meniskus");
|
|
}
|
|
|
|
#[tokio::test]
|
|
async fn ollama_oneliner_returns_status_on_500() {
|
|
let server = MockServer::start().await;
|
|
Mock::given(method("POST"))
|
|
.and(path("/api/chat"))
|
|
.respond_with(ResponseTemplate::new(500).set_body_string("boom"))
|
|
.mount(&server)
|
|
.await;
|
|
|
|
let client = reqwest::Client::new();
|
|
let err = generate_oneliner(&client, &server.uri(), MODEL, 0, "x", Duration::from_secs(5))
|
|
.await
|
|
.expect_err("expected error");
|
|
match err {
|
|
OllamaError::Status { status, body } => {
|
|
assert_eq!(status, 500);
|
|
assert_eq!(body, "boom");
|
|
}
|
|
other => panic!("unexpected: {other:?}"),
|
|
}
|
|
}
|
|
|
|
#[tokio::test]
|
|
async fn ollama_oneliner_sends_keep_alive_and_model() {
|
|
let server = MockServer::start().await;
|
|
Mock::given(method("POST"))
|
|
.and(path("/api/chat"))
|
|
.and(body_partial_json(json!({
|
|
"model": MODEL,
|
|
"stream": false,
|
|
"keep_alive": 0
|
|
})))
|
|
.respond_with(ResponseTemplate::new(200).set_body_json(ok_chat_body("OK")))
|
|
.expect(1)
|
|
.mount(&server)
|
|
.await;
|
|
|
|
let client = reqwest::Client::new();
|
|
let _ = generate_oneliner(&client, &server.uri(), MODEL, 0, "irrelevant", Duration::from_secs(5))
|
|
.await
|
|
.expect("call failed");
|
|
// Mock::expect(1) + server drop verifies the matcher hit exactly once.
|
|
}
|
|
|
|
#[tokio::test]
|
|
async fn ollama_oneliner_parse_error_on_missing_content() {
|
|
let server = MockServer::start().await;
|
|
Mock::given(method("POST"))
|
|
.and(path("/api/chat"))
|
|
.respond_with(ResponseTemplate::new(200).set_body_json(json!({"done": true})))
|
|
.mount(&server)
|
|
.await;
|
|
|
|
let client = reqwest::Client::new();
|
|
let err = generate_oneliner(&client, &server.uri(), MODEL, 0, "x", Duration::from_secs(5))
|
|
.await
|
|
.expect_err("expected parse error");
|
|
assert!(matches!(err, OllamaError::Parse(_)), "got {err:?}");
|
|
}
|