use super::RealtimeTranscriber;
use crate::config::MistralConfig;
use crate::error::TalkError;
use async_trait::async_trait;
use base64::engine::general_purpose::STANDARD as BASE64_STANDARD;
use base64::Engine;
use futures::stream::SplitSink;
use futures::{SinkExt, StreamExt};
use std::collections::{BTreeMap, BTreeSet};
use std::time::Duration;
use tokio::sync::mpsc;
use tokio_tungstenite::tungstenite::Message;
use tokio_util::sync::CancellationToken;
const DEFAULT_REALTIME_MODEL: &str = "voxtral-mini-transcribe-realtime-2602";
const DEFAULT_REALTIME_ENDPOINT: &str = "wss://api.mistral.ai";
const REALTIME_PATH: &str = "/v1/audio/transcriptions/realtime";
const SESSION_CREATED_TIMEOUT: Duration = Duration::from_secs(15);
const WS_PING_INTERVAL: Duration = Duration::from_secs(30);
#[derive(Debug, Clone)]
pub enum TranscriptionEvent {
SessionCreated,
TextDelta { text: String },
SegmentDelta {
text: String,
start: Option<f64>,
end: Option<f64>,
},
ItemCreated {
item_id: String,
previous_item_id: Option<String>,
},
ItemTextDelta {
item_id: String,
content_index: u64,
text: String,
},
ItemTextCompleted {
item_id: String,
content_index: u64,
transcript: String,
},
Language { language: String },
SessionInfo {
session_id: Option<String>,
conversation_id: Option<String>,
},
RateLimitsUpdated { raw: serde_json::Value },
TransportMetadata { headers: BTreeMap<String, String> },
Done,
Error { message: String },
Unknown {
event_type: Option<String>,
raw: String,
},
}
#[derive(Debug, Default)]
struct ItemText {
provisional: String,
completed: Option<String>,
emitted: bool,
}
#[derive(Debug)]
struct ItemOrder {
previous_item_id: Option<String>,
first_seen: usize,
order_known: bool,
}
#[derive(Debug, Default)]
pub struct OrderedItemTranscript {
text: BTreeMap<(String, u64), ItemText>,
order: BTreeMap<String, ItemOrder>,
next_first_seen: usize,
}
impl OrderedItemTranscript {
pub fn item_created(&mut self, item_id: &str, previous_item_id: Option<&str>) {
self.ensure_item(item_id);
if let Some(order) = self.order.get_mut(item_id) {
order.order_known = true;
if let Some(previous_item_id) = previous_item_id {
order.previous_item_id = Some(previous_item_id.to_string());
}
}
}
pub fn append_delta(&mut self, item_id: &str, content_index: u64, text: &str) {
self.ensure_item(item_id);
self.text
.entry((item_id.to_string(), content_index))
.or_default()
.provisional
.push_str(text);
}
pub fn complete(&mut self, item_id: &str, content_index: u64, transcript: &str) {
self.ensure_item(item_id);
self.text
.entry((item_id.to_string(), content_index))
.or_default()
.completed = Some(transcript.to_string());
}
pub fn snapshot(&self) -> String {
self.render(false).join(" ")
}
pub fn drain_completed_prefix(&mut self) -> Vec<String> {
let mut drained = Vec::new();
for item_id in self.ordered_item_ids() {
let Some(order) = self.order.get(&item_id) else {
break;
};
if !order.order_known {
break;
}
if let Some(previous_item_id) = order.previous_item_id.as_deref() {
if !self.item_is_resolved(previous_item_id) {
break;
}
}
let keys = self.item_keys(&item_id);
if keys.is_empty() {
break;
}
let mut expected_index = 0;
for key in keys {
if key.1 != expected_index {
return drained;
}
let Some(item_text) = self.text.get_mut(&key) else {
return drained;
};
if item_text.emitted {
expected_index += 1;
continue;
}
let Some(completed) = item_text.completed.as_deref() else {
return drained;
};
if !completed.is_empty() {
drained.push(completed.to_string());
}
item_text.emitted = true;
expected_index += 1;
}
}
drained
}
pub fn drain_terminal(&mut self) -> Vec<String> {
let mut drained = self.drain_completed_prefix();
for item_id in self.ordered_item_ids() {
for key in self.item_keys(&item_id) {
let Some(item_text) = self.text.get_mut(&key) else {
continue;
};
if item_text.emitted {
continue;
}
let text = item_text
.completed
.as_deref()
.unwrap_or(item_text.provisional.as_str());
if !text.is_empty() {
drained.push(text.to_string());
}
item_text.emitted = true;
}
}
drained
}
pub(crate) fn reset_generation(&mut self) {
self.text.clear();
self.order.clear();
self.next_first_seen = 0;
}
pub(crate) fn is_empty(&self) -> bool {
self.text.is_empty()
}
fn ensure_item(&mut self, item_id: &str) {
if self.order.contains_key(item_id) {
return;
}
self.order.insert(
item_id.to_string(),
ItemOrder {
previous_item_id: None,
first_seen: self.next_first_seen,
order_known: false,
},
);
self.next_first_seen += 1;
}
fn render(&self, completed_only: bool) -> Vec<String> {
self.ordered_item_ids()
.into_iter()
.filter_map(|item_id| {
let parts = self
.text
.iter()
.filter(|((id, _), _)| id == &item_id)
.filter_map(|(_, item_text)| {
if completed_only {
item_text.completed.as_deref()
} else {
item_text
.completed
.as_deref()
.or(Some(item_text.provisional.as_str()))
}
})
.filter(|part| !part.is_empty())
.collect::<Vec<_>>()
.join(" ");
(!parts.is_empty()).then_some(parts)
})
.collect()
}
fn item_keys(&self, item_id: &str) -> Vec<(String, u64)> {
self.text
.keys()
.filter(|(id, _)| id == item_id)
.cloned()
.collect()
}
fn item_is_resolved(&self, item_id: &str) -> bool {
let keys = self.item_keys(item_id);
!keys.is_empty()
&& keys
.iter()
.enumerate()
.all(|(index, key)| key.1 == index as u64 && self.text[key].emitted)
}
fn ordered_item_ids(&self) -> Vec<String> {
let mut first_seen = self.order.keys().cloned().collect::<Vec<_>>();
first_seen.sort_by_key(|item_id| self.order[item_id].first_seen);
let mut visiting = BTreeSet::new();
let mut emitted = BTreeSet::new();
let mut ordered = Vec::new();
for item_id in first_seen {
self.visit_item(&item_id, &mut visiting, &mut emitted, &mut ordered);
}
ordered
}
fn visit_item(
&self,
item_id: &str,
visiting: &mut BTreeSet<String>,
emitted: &mut BTreeSet<String>,
ordered: &mut Vec<String>,
) {
if emitted.contains(item_id) || !visiting.insert(item_id.to_string()) {
return;
}
if let Some(previous_item_id) = self
.order
.get(item_id)
.and_then(|item| item.previous_item_id.as_deref())
{
if self.order.contains_key(previous_item_id) {
self.visit_item(previous_item_id, visiting, emitted, ordered);
}
}
visiting.remove(item_id);
if emitted.insert(item_id.to_string()) {
ordered.push(item_id.to_string());
}
}
}
pub fn parse_event(json_str: &str) -> TranscriptionEvent {
let value: serde_json::Value = match serde_json::from_str(json_str) {
Ok(v) => v,
Err(_) => {
return TranscriptionEvent::Unknown {
event_type: None,
raw: json_str.to_string(),
};
}
};
let event_type = value.get("type").and_then(|v| v.as_str());
match event_type {
Some("session.created") | Some("session.updated") => TranscriptionEvent::SessionCreated,
Some("transcription.text.delta") => {
let text = value
.get("text")
.and_then(|v| v.as_str())
.unwrap_or("")
.to_string();
TranscriptionEvent::TextDelta { text }
}
Some("transcription.segment") => {
let text = value
.get("text")
.and_then(|v| v.as_str())
.unwrap_or("")
.to_string();
let start = value.get("start").and_then(|v| v.as_f64());
let end = value.get("end").and_then(|v| v.as_f64());
TranscriptionEvent::SegmentDelta { text, start, end }
}
Some("transcription.language") => {
let language = value
.get("language")
.and_then(|v| v.as_str())
.unwrap_or("")
.to_string();
TranscriptionEvent::Language { language }
}
Some("transcription.done") => TranscriptionEvent::Done,
Some("error") => {
let message = value
.get("error")
.and_then(|v| v.get("message"))
.and_then(|v| v.as_str())
.unwrap_or("unknown error")
.to_string();
TranscriptionEvent::Error { message }
}
_ => TranscriptionEvent::Unknown {
event_type: event_type.map(|s| s.to_string()),
raw: json_str.to_string(),
},
}
}
pub fn build_ws_url(endpoint: &str, model: &str) -> String {
format!("{}{}?model={}", endpoint, REALTIME_PATH, model)
}
pub fn pcm_to_bytes(samples: &[i16]) -> Vec<u8> {
let mut bytes = Vec::with_capacity(samples.len() * 2);
for &sample in samples {
bytes.extend_from_slice(&sample.to_le_bytes());
}
bytes
}
pub fn pcm_bytes_to_base64(bytes: &[u8]) -> String {
BASE64_STANDARD.encode(bytes)
}
fn http_to_ws(url: &str) -> String {
if let Some(rest) = url.strip_prefix("https://") {
format!("wss://{}", rest)
} else if let Some(rest) = url.strip_prefix("http://") {
format!("ws://{}", rest)
} else {
url.to_string()
}
}
pub struct MistralRealtimeTranscriber {
config: MistralConfig,
endpoint: String,
sink: std::sync::Arc<dyn crate::telemetry::TelemetrySink>,
cancel_token: CancellationToken,
}
impl MistralRealtimeTranscriber {
pub fn new(config: MistralConfig) -> Self {
let endpoint = config
.url
.as_deref()
.map(|u| http_to_ws(u.trim_end_matches('/')))
.unwrap_or_else(|| DEFAULT_REALTIME_ENDPOINT.to_string());
Self {
config,
endpoint,
sink: std::sync::Arc::new(crate::telemetry::NoOpSink),
cancel_token: CancellationToken::new(),
}
}
#[cfg(test)]
pub fn with_endpoint(config: MistralConfig, endpoint: String) -> Self {
Self {
config,
endpoint,
sink: std::sync::Arc::new(crate::telemetry::NoOpSink),
cancel_token: CancellationToken::new(),
}
}
fn realtime_model(&self) -> &str {
DEFAULT_REALTIME_MODEL
}
async fn validate_realtime_session(&self) -> Result<(), TalkError> {
let ws_url = build_ws_url(&self.endpoint, self.realtime_model());
log::debug!("validation: connecting to {}", ws_url);
let req = super::transport::Request {
method: super::transport::Method::Get,
url: ws_url.clone(),
headers: vec![(
"Authorization".into(),
format!("Bearer {}", self.config.api_key),
)],
body: super::transport::RequestBody::Empty,
provider: crate::config::Provider::Mistral,
provider_name: "Mistral".into(),
phase: crate::error::PipelinePhase::Validate,
wall_clock: None,
};
let ws_stream = super::transport::ws_upgrade(req, &self.sink, self.cancel_token.clone())
.await
.map_err(|pf| TalkError::Config(pf.to_string()))?;
let (mut sink_split, mut source) = ws_stream.split();
let result = tokio::time::timeout(
SESSION_CREATED_TIMEOUT,
wait_for_session_created(&mut source),
)
.await
.map_err(|_| {
TalkError::Config(format!(
"Timed out waiting for session.created after {}s",
SESSION_CREATED_TIMEOUT.as_secs()
))
})?;
let _ = sink_split.send(Message::Close(None)).await;
result.map(|_| ())
}
pub async fn transcribe_realtime(
&self,
audio_rx: mpsc::Receiver<Vec<i16>>,
) -> Result<mpsc::Receiver<TranscriptionEvent>, TalkError> {
let ws_url = build_ws_url(&self.endpoint, self.realtime_model());
log::debug!("connecting to WebSocket: {}", ws_url);
let req = super::transport::Request {
method: super::transport::Method::Get,
url: ws_url.clone(),
headers: vec![(
"Authorization".into(),
format!("Bearer {}", self.config.api_key),
)],
body: super::transport::RequestBody::Empty,
provider: crate::config::Provider::Mistral,
provider_name: "Mistral".into(),
phase: crate::error::PipelinePhase::Request,
wall_clock: None,
};
let ws_stream = super::transport::ws_upgrade(req, &self.sink, self.cancel_token.clone())
.await
.map_err(|pf| TalkError::Transcription(pf.to_string()))?;
let (mut ws_sink, mut ws_source) = ws_stream.split();
self.sink
.emit(crate::telemetry::TranscriptionEvent::Status {
message: "session handshake…".into(),
t: std::time::Instant::now(),
});
log::debug!("WebSocket connected, waiting for session.created");
let session_event = tokio::time::timeout(
SESSION_CREATED_TIMEOUT,
wait_for_session_created(&mut ws_source),
)
.await
.map_err(|_| {
TalkError::Transcription(format!(
"Timed out waiting for session.created after {}s",
SESSION_CREATED_TIMEOUT.as_secs()
))
})??;
log::info!("realtime session established");
self.sink
.emit(crate::telemetry::TranscriptionEvent::Status {
message: "session ready, awaiting audio…".into(),
t: std::time::Instant::now(),
});
log::debug!("sending session.update with pcm_s16le/16000");
let session_update = serde_json::json!({
"type": "session.update",
"session": {
"audio_format": {
"encoding": "pcm_s16le",
"sample_rate": 16000
}
}
});
ws_sink
.send(Message::Text(session_update.to_string()))
.await
.map_err(|e| {
TalkError::Transcription(format!("Failed to send session.update: {}", e))
})?;
self.sink
.emit(crate::telemetry::TranscriptionEvent::Status {
message: "streaming audio…".into(),
t: std::time::Instant::now(),
});
let (event_tx, event_rx) = mpsc::channel::<TranscriptionEvent>(100);
let _ = event_tx.send(session_event).await;
let cancel = CancellationToken::new();
let sender_task = tokio::spawn(sender_loop(audio_rx, ws_sink, cancel.clone()));
let receiver_task = tokio::spawn(receiver_loop(ws_source, event_tx, cancel.clone()));
tokio::spawn(async move {
if let Err(e) = sender_task.await {
log::error!("sender task panicked: {}", e);
}
if let Err(e) = receiver_task.await {
log::error!("receiver task panicked: {}", e);
}
});
Ok(event_rx)
}
}
#[async_trait]
impl RealtimeTranscriber for MistralRealtimeTranscriber {
async fn validate(&self) -> Result<(), TalkError> {
self.validate_realtime_session().await
}
async fn transcribe_realtime(
&self,
audio_rx: mpsc::Receiver<Vec<i16>>,
) -> Result<mpsc::Receiver<TranscriptionEvent>, TalkError> {
self.transcribe_realtime(audio_rx).await
}
fn set_sink(&mut self, sink: std::sync::Arc<dyn crate::telemetry::TelemetrySink>) {
self.sink = sink;
}
fn set_cancel_token(&mut self, token: CancellationToken) {
self.cancel_token = token;
}
}
async fn wait_for_session_created<S>(ws_source: &mut S) -> Result<TranscriptionEvent, TalkError>
where
S: futures::Stream<Item = Result<Message, tokio_tungstenite::tungstenite::Error>> + Unpin,
{
while let Some(msg_result) = ws_source.next().await {
let msg = msg_result.map_err(|e| {
TalkError::Transcription(format!(
"WebSocket error waiting for session.created: {}",
e
))
})?;
if let Message::Text(text) = msg {
let event = parse_event(&text);
match event {
TranscriptionEvent::SessionCreated => return Ok(event),
TranscriptionEvent::Error { ref message } => {
return Err(TalkError::Transcription(format!(
"Server error during session setup: {}",
message
)));
}
_ => {
}
}
}
}
Err(TalkError::Transcription(
"WebSocket closed before session.created received".to_string(),
))
}
async fn sender_loop<S>(
mut audio_rx: mpsc::Receiver<Vec<i16>>,
mut ws_sink: SplitSink<S, Message>,
cancel: CancellationToken,
) where
S: futures::Sink<Message> + Unpin,
<S as futures::Sink<Message>>::Error: std::fmt::Display,
{
let mut ping_interval = tokio::time::interval(WS_PING_INTERVAL);
ping_interval.tick().await;
loop {
tokio::select! {
chunk = audio_rx.recv() => {
match chunk {
Some(pcm_chunk) => {
let bytes = pcm_to_bytes(&pcm_chunk);
log::trace!("sending audio chunk: {} PCM samples, {} bytes", pcm_chunk.len(), bytes.len());
let b64 = pcm_bytes_to_base64(&bytes);
let msg = serde_json::json!({
"type": "input_audio.append",
"audio": b64
});
if let Err(e) = ws_sink.send(Message::Text(msg.to_string())).await {
log::error!("WebSocket send error: {}", e);
cancel.cancel();
return;
}
}
None => break, }
}
_ = ping_interval.tick() => {
if let Err(e) = ws_sink.send(Message::Ping(vec![])).await {
log::warn!("WebSocket ping failed: {}", e);
cancel.cancel();
return;
}
}
_ = cancel.cancelled() => {
return;
}
}
}
let end_msg = serde_json::json!({
"type": "input_audio.end"
});
log::debug!("sending input_audio.end");
if let Err(e) = ws_sink.send(Message::Text(end_msg.to_string())).await {
log::error!("WebSocket send error (input_audio.end): {}", e);
cancel.cancel();
}
}
async fn receiver_loop<S>(
mut ws_source: S,
event_tx: mpsc::Sender<TranscriptionEvent>,
cancel: CancellationToken,
) where
S: futures::Stream<Item = Result<Message, tokio_tungstenite::tungstenite::Error>> + Unpin,
{
loop {
tokio::select! {
msg_opt = ws_source.next() => {
let msg_result = match msg_opt {
Some(r) => r,
None => {
log::warn!("WebSocket stream ended unexpectedly");
cancel.cancel();
return;
}
};
let msg = match msg_result {
Ok(m) => m,
Err(e) => {
log::error!("WebSocket receive error: {}", e);
let _ = event_tx
.send(TranscriptionEvent::Error {
message: format!("WebSocket error: {}", e),
})
.await;
cancel.cancel();
return;
}
};
match msg {
Message::Text(text) => {
log::trace!("received WS text: {}", text);
let event = parse_event(&text);
let is_terminal = matches!(
event,
TranscriptionEvent::Done | TranscriptionEvent::Error { .. }
);
if event_tx.send(event).await.is_err() {
cancel.cancel();
return;
}
if is_terminal {
return;
}
}
Message::Close(frame) => {
log::debug!("received WS Close frame: {:?}", frame);
let _ = event_tx.send(TranscriptionEvent::Done).await;
return;
}
Message::Pong(_) => {
log::trace!("received WS Pong");
}
_ => {
}
}
}
_ = cancel.cancelled() => {
return;
}
}
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_http_to_ws_https() {
assert_eq!(http_to_ws("https://api.mistral.ai"), "wss://api.mistral.ai");
}
#[test]
fn test_http_to_ws_http() {
assert_eq!(http_to_ws("http://localhost:8080"), "ws://localhost:8080");
}
#[test]
fn test_http_to_ws_already_wss() {
assert_eq!(http_to_ws("wss://api.mistral.ai"), "wss://api.mistral.ai");
}
#[test]
fn test_http_to_ws_already_ws() {
assert_eq!(http_to_ws("ws://localhost:8080"), "ws://localhost:8080");
}
#[test]
fn test_new_uses_custom_url() {
let config = MistralConfig {
api_key: "key".to_string(),
url: Some("https://custom.example.com".to_string()),
model: "voxtral-mini-2507".to_string(),
context_bias: None,
tts_model: "voxtral-mini-tts-latest".to_string(),
tts_voice: None,
tts_voices: None,
};
let transcriber = MistralRealtimeTranscriber::new(config);
assert_eq!(transcriber.endpoint, "wss://custom.example.com");
}
#[test]
fn test_new_default_endpoint() {
let config = MistralConfig {
api_key: "key".to_string(),
url: None,
model: "voxtral-mini-2507".to_string(),
context_bias: None,
tts_model: "voxtral-mini-tts-latest".to_string(),
tts_voice: None,
tts_voices: None,
};
let transcriber = MistralRealtimeTranscriber::new(config);
assert_eq!(transcriber.endpoint, "wss://api.mistral.ai");
}
#[test]
fn test_parse_event_text_delta() {
let json = r#"{"type": "transcription.text.delta", "text": "hello"}"#;
let event = parse_event(json);
match event {
TranscriptionEvent::TextDelta { text } => assert_eq!(text, "hello"),
other => panic!("Expected TextDelta, got {:?}", other),
}
}
#[test]
fn test_parse_event_done() {
let json = r#"{"type": "transcription.done"}"#;
let event = parse_event(json);
assert!(matches!(event, TranscriptionEvent::Done));
}
#[test]
fn test_parse_event_error() {
let json = r#"{"type": "error", "error": {"message": "bad"}}"#;
let event = parse_event(json);
match event {
TranscriptionEvent::Error { message } => assert_eq!(message, "bad"),
other => panic!("Expected Error, got {:?}", other),
}
}
#[test]
fn test_parse_event_session_created() {
let json = r#"{"type": "session.created", "session": {}}"#;
let event = parse_event(json);
assert!(matches!(event, TranscriptionEvent::SessionCreated));
}
#[test]
fn test_parse_event_language() {
let json = r#"{"type": "transcription.language", "language": "en"}"#;
let event = parse_event(json);
match event {
TranscriptionEvent::Language { language } => assert_eq!(language, "en"),
other => panic!("Expected Language, got {:?}", other),
}
}
#[test]
fn test_parse_event_segment() {
let json =
r#"{"type": "transcription.segment", "text": "hello world", "start": 0.0, "end": 1.5}"#;
let event = parse_event(json);
match event {
TranscriptionEvent::SegmentDelta { text, start, end } => {
assert_eq!(text, "hello world");
assert_eq!(start, Some(0.0));
assert_eq!(end, Some(1.5));
}
other => panic!("Expected SegmentDelta, got {:?}", other),
}
}
#[test]
fn ordered_items_replace_provisional_with_completed() {
let mut transcript = OrderedItemTranscript::default();
transcript.item_created("item-1", None);
transcript.append_delta("item-1", 0, "Hello");
transcript.append_delta("item-1", 0, " world");
assert_eq!(transcript.snapshot(), "Hello world");
transcript.complete("item-1", 0, "Hello, corrected world.");
assert_eq!(transcript.snapshot(), "Hello, corrected world.");
assert_eq!(transcript.snapshot(), "Hello, corrected world.");
}
#[test]
fn ordered_items_render_conversation_order_when_completed_in_reverse() {
let mut transcript = OrderedItemTranscript::default();
transcript.item_created("item-1", None);
transcript.item_created("item-2", Some("item-1"));
transcript.complete("item-2", 0, "second");
assert!(transcript.drain_completed_prefix().is_empty());
transcript.complete("item-1", 0, "first");
assert_eq!(transcript.snapshot(), "first second");
assert_eq!(
transcript.drain_completed_prefix(),
vec!["first".to_string(), "second".to_string()]
);
}
#[test]
fn ordered_items_late_order_metadata_unblocks_completed_prefix() {
let mut transcript = OrderedItemTranscript::default();
transcript.complete("item-2", 0, "second");
assert!(transcript.drain_completed_prefix().is_empty());
transcript.item_created("item-1", None);
transcript.complete("item-1", 0, "first");
transcript.item_created("item-2", Some("item-1"));
assert_eq!(
transcript.drain_completed_prefix(),
vec!["first".to_string(), "second".to_string()]
);
}
#[test]
fn ordered_items_drain_multiple_content_indexes_in_order() {
let mut transcript = OrderedItemTranscript::default();
transcript.item_created("item-1", None);
transcript.append_delta("item-1", 0, "zero provisional");
transcript.append_delta("item-1", 1, "one provisional");
transcript.complete("item-1", 1, "one");
assert!(transcript.drain_completed_prefix().is_empty());
transcript.complete("item-1", 0, "zero");
assert_eq!(
transcript.drain_completed_prefix(),
vec!["zero".to_string(), "one".to_string()]
);
}
#[test]
fn ordered_items_terminal_drain_preserves_provisional_once() {
let mut transcript = OrderedItemTranscript::default();
transcript.item_created("item-1", None);
transcript.append_delta("item-1", 0, "provisional terminal text");
assert_eq!(
transcript.drain_terminal(),
vec!["provisional terminal text".to_string()]
);
assert!(transcript.drain_terminal().is_empty());
}
#[test]
fn ordered_items_empty_completion_replaces_provisional_without_emission() {
let mut transcript = OrderedItemTranscript::default();
transcript.item_created("item-1", None);
transcript.append_delta("item-1", 0, "discard me");
transcript.complete("item-1", 0, "");
assert!(transcript.drain_completed_prefix().is_empty());
assert!(transcript.drain_terminal().is_empty());
assert_eq!(transcript.snapshot(), "");
}
#[test]
fn ordered_items_accept_completion_without_delta() {
let mut transcript = OrderedItemTranscript::default();
transcript.complete("item-1", 0, "authoritative only");
assert_eq!(transcript.snapshot(), "authoritative only");
assert_eq!(transcript.snapshot(), "authoritative only");
}
#[test]
fn ordered_items_fall_back_to_first_seen_order_without_links() {
let mut transcript = OrderedItemTranscript::default();
transcript.complete("item-b", 0, "first seen");
transcript.complete("item-a", 0, "second seen");
assert_eq!(transcript.snapshot(), "first seen second seen");
}
#[test]
fn ordered_items_generation_reset_clears_active_item_ids() {
let mut transcript = OrderedItemTranscript::default();
transcript.item_created("old-item", None);
transcript.complete("old-item", 0, "old");
assert_eq!(transcript.drain_completed_prefix(), vec!["old".to_string()]);
transcript.reset_generation();
transcript.item_created("fresh-item", None);
transcript.complete("fresh-item", 0, "fresh");
assert_eq!(transcript.snapshot(), "fresh");
assert_eq!(
transcript.drain_completed_prefix(),
vec!["fresh".to_string()]
);
}
#[test]
fn test_parse_event_unknown_type() {
let json = r#"{"type": "foo"}"#;
let event = parse_event(json);
match event {
TranscriptionEvent::Unknown { event_type, raw } => {
assert_eq!(event_type, Some("foo".to_string()));
assert_eq!(raw, json);
}
other => panic!("Expected Unknown, got {:?}", other),
}
}
#[test]
fn test_parse_event_invalid_json() {
let json = "not json{{";
let event = parse_event(json);
match event {
TranscriptionEvent::Unknown { event_type, raw } => {
assert!(event_type.is_none());
assert_eq!(raw, json);
}
other => panic!("Expected Unknown, got {:?}", other),
}
}
#[test]
fn test_build_ws_url() {
let url = build_ws_url("wss://api.mistral.ai", "my-model");
assert_eq!(
url,
"wss://api.mistral.ai/v1/audio/transcriptions/realtime?model=my-model"
);
}
#[test]
fn test_build_ws_url_custom_endpoint() {
let url = build_ws_url("wss://custom.example.com", "test-model");
assert_eq!(
url,
"wss://custom.example.com/v1/audio/transcriptions/realtime?model=test-model"
);
}
#[test]
fn test_pcm_to_base64() {
let samples: Vec<i16> = vec![256, 32767];
let bytes = pcm_to_bytes(&samples);
assert_eq!(bytes, vec![0x00, 0x01, 0xFF, 0x7F]);
let b64 = pcm_bytes_to_base64(&bytes);
let decoded = BASE64_STANDARD.decode(&b64).expect("valid base64");
assert_eq!(decoded, bytes);
}
#[test]
fn test_pcm_to_bytes_empty() {
let samples: Vec<i16> = vec![];
let bytes = pcm_to_bytes(&samples);
assert!(bytes.is_empty());
}
#[test]
fn test_pcm_to_bytes_negative() {
let samples: Vec<i16> = vec![-1, -32768];
let bytes = pcm_to_bytes(&samples);
assert_eq!(bytes, vec![0xFF, 0xFF, 0x00, 0x80]);
}
#[test]
fn test_parse_event_error_missing_message() {
let json = r#"{"type": "error", "error": {}}"#;
let event = parse_event(json);
match event {
TranscriptionEvent::Error { message } => {
assert_eq!(message, "unknown error");
}
other => panic!("Expected Error, got {:?}", other),
}
}
#[test]
fn test_parse_event_text_delta_missing_text() {
let json = r#"{"type": "transcription.text.delta"}"#;
let event = parse_event(json);
match event {
TranscriptionEvent::TextDelta { text } => assert_eq!(text, ""),
other => panic!("Expected TextDelta, got {:?}", other),
}
}
#[test]
fn test_timeout_constants_are_reasonable() {
const _: () = {
const SUM: u64 = 2 + 5 + 8 + 11 + 15;
assert!(SUM >= 5);
assert!(SUM <= 120);
};
assert!(SESSION_CREATED_TIMEOUT.as_secs() >= 5);
assert!(SESSION_CREATED_TIMEOUT.as_secs() <= 60);
assert!(WS_PING_INTERVAL.as_secs() >= 10);
assert!(WS_PING_INTERVAL.as_secs() <= 120);
}
#[tokio::test]
async fn validate_fails_on_rejected_ws_upgrade_without_streaming() {
use wiremock::matchers::{method, path};
use wiremock::{Mock, MockServer, ResponseTemplate};
let server = MockServer::start().await;
let upgrade = Mock::given(method("GET"))
.and(path(REALTIME_PATH))
.respond_with(ResponseTemplate::new(401).set_body_string("Unauthorized"))
.expect(1)
.mount_as_scoped(&server)
.await;
let config = MistralConfig {
api_key: "bad-key".to_string(),
url: None,
model: "voxtral-mini-transcribe-realtime-2602".to_string(),
context_bias: None,
tts_model: "voxtral-mini-tts-latest".to_string(),
tts_voice: None,
tts_voices: None,
};
let transcriber =
MistralRealtimeTranscriber::with_endpoint(config, http_to_ws(&server.uri()));
let result = transcriber.validate().await;
let err = match result {
Err(e) => e.to_string(),
Ok(()) => panic!("validate must fail when the upgrade is rejected"),
};
assert!(
err.contains("401"),
"error must carry the HTTP status: {err}"
);
drop(upgrade); }
}