fix(voice): bound dictation WS audio buffer (partial fix for #1924) (#2238)

This commit is contained in:
sanil-23
2026-05-20 01:59:58 +05:30
committed by GitHub
parent b8fbb364fb
commit cf5dd2d7c7
+146 -11
View File
@@ -12,6 +12,20 @@
//! `{"type":"stop"}` text frame
//! { "type": "error", "message": "..." } — on error
//! Client → Server: text frame `{"type":"stop"}` — end recording, get final result
//!
//! # Security notes
//!
//! ## Authentication
//! `GET /ws/dictation` is intentionally exempt from Bearer-token authentication because
//! the browser WebSocket API cannot set arbitrary request headers on upgrade. The correct
//! auth mechanism is a separate maintainer decision; see `src/core/auth.rs` for the
//! documented exemption. Do NOT add a Bearer-header check here — it will not work from
//! browsers and the design decision is tracked in issue #1924.
//!
//! ## Memory cap
//! The full-audio accumulation buffer (`full_audio_buf`) is bounded by
//! `MAX_FULL_AUDIO_SAMPLES` (~5 min at 16 kHz). Clients that stream beyond this limit
//! are disconnected with an error frame; see `append_stream_samples`.
use std::sync::atomic::{AtomicU64, Ordering};
use std::sync::Arc;
@@ -31,6 +45,13 @@ const AUDIO_SAMPLE_RATE: usize = 16_000;
const MIN_PARTIAL_SAMPLES: usize = AUDIO_SAMPLE_RATE / 2; // 0.5s
const MAX_STREAM_BUFFER_SAMPLES: usize = AUDIO_SAMPLE_RATE * 15; // 15s sliding window
/// Hard cap on the full-audio accumulation buffer.
///
/// Derived from `AUDIO_SAMPLE_RATE` (16 kHz mono PCM16) × 60 s × 5 min = 4 800 000 samples
/// ≈ 9.6 MiB per connection. Clients that send audio beyond this limit are disconnected
/// gracefully with a `{"type":"error"}` frame so the server never OOMs (issue #1924).
const MAX_FULL_AUDIO_SAMPLES: usize = AUDIO_SAMPLE_RATE * 60 * 5; // ~5 minutes
#[derive(Debug, Deserialize)]
struct ClientCommand {
#[serde(rename = "type")]
@@ -49,7 +70,30 @@ fn decode_pcm16le_frame(data: &[u8]) -> Option<Vec<i16>> {
)
}
fn append_stream_samples(audio_buf: &mut Vec<i16>, full_audio_buf: &mut Vec<i16>, samples: &[i16]) {
/// Append `samples` to both the sliding window buffer and the full-audio accumulation
/// buffer, enforcing the hard cap on the latter.
///
/// Returns `true` when the full-audio buffer is within the allowed limit (normal path).
/// Returns `false` when appending `samples` would push `full_audio_buf` beyond
/// `MAX_FULL_AUDIO_SAMPLES`; in that case the samples are **not** appended and the caller
/// must disconnect the client to prevent unbounded memory growth (issue #1924).
fn append_stream_samples(
audio_buf: &mut Vec<i16>,
full_audio_buf: &mut Vec<i16>,
samples: &[i16],
) -> bool {
// Enforce hard cap on the full-audio accumulation buffer first.
if full_audio_buf.len().saturating_add(samples.len()) > MAX_FULL_AUDIO_SAMPLES {
log::warn!(
"{LOG_PREFIX} full_audio_buf cap reached ({} / {} samples); refusing to append {} \
more samples — client will be disconnected",
full_audio_buf.len(),
MAX_FULL_AUDIO_SAMPLES,
samples.len(),
);
return false;
}
full_audio_buf.extend_from_slice(samples);
audio_buf.extend_from_slice(samples);
if audio_buf.len() > MAX_STREAM_BUFFER_SAMPLES {
@@ -61,6 +105,7 @@ fn append_stream_samples(audio_buf: &mut Vec<i16>, full_audio_buf: &mut Vec<i16>
audio_buf.len()
);
}
true
}
fn is_stop_command(text: &str) -> bool {
@@ -158,15 +203,44 @@ pub async fn handle_dictation_ws(mut socket: WebSocket, config: Arc<Config>) {
continue;
};
let mut full = full_audio_buf.lock().await;
let mut buf = audio_buf.lock().await;
append_stream_samples(&mut buf, &mut full, &samples);
audio_revision.fetch_add(1, Ordering::Relaxed);
log::trace!(
"{LOG_PREFIX} buffered {} new samples, total {}",
samples.len(),
buf.len()
);
let cap_exceeded = {
let mut full = full_audio_buf.lock().await;
let mut buf = audio_buf.lock().await;
let ok = append_stream_samples(&mut buf, &mut full, &samples);
if ok {
audio_revision.fetch_add(1, Ordering::Relaxed);
log::trace!(
"{LOG_PREFIX} buffered {} new samples, total {}",
samples.len(),
buf.len()
);
}
!ok
};
if cap_exceeded {
// Send an error frame and close — never OOM.
let err_msg = serde_json::json!({
"type": "error",
"message": format!(
"Recording limit reached: maximum {} minutes of audio per session",
MAX_FULL_AUDIO_SAMPLES / AUDIO_SAMPLE_RATE / 60
),
});
let _ = socket
.send(Message::Text(err_msg.to_string().into()))
.await;
log::warn!(
"{LOG_PREFIX} disconnecting client: full_audio_buf cap ({} samples, \
{} min at 16 kHz) exceeded",
MAX_FULL_AUDIO_SAMPLES,
MAX_FULL_AUDIO_SAMPLES / AUDIO_SAMPLE_RATE / 60,
);
if let Some(h) = inference_handle {
h.abort();
}
return;
}
}
Some(Ok(Message::Text(text))) => {
@@ -272,13 +346,74 @@ mod tests {
fn append_stream_samples_keeps_full_audio_and_trims_window() {
let mut audio = vec![0; MAX_STREAM_BUFFER_SAMPLES - 2];
let mut full = vec![1, 2];
append_stream_samples(&mut audio, &mut full, &[3, 4, 5, 6]);
let ok = append_stream_samples(&mut audio, &mut full, &[3, 4, 5, 6]);
assert!(ok, "should succeed when under cap");
assert_eq!(full, vec![1, 2, 3, 4, 5, 6]);
assert_eq!(audio.len(), MAX_STREAM_BUFFER_SAMPLES);
assert_eq!(&audio[audio.len() - 4..], &[3, 4, 5, 6]);
}
/// Feed enough samples to hit the full-audio cap and verify:
/// 1. The buffer does NOT grow past `MAX_FULL_AUDIO_SAMPLES`.
/// 2. `append_stream_samples` returns `false` (cap-exceeded signal) when the next
/// chunk would overflow.
#[test]
fn append_stream_samples_enforces_full_audio_cap() {
let chunk_size = 1_024usize;
let mut audio = Vec::new();
let mut full = Vec::new();
// Fill up to exactly the cap in chunks.
let full_chunks = MAX_FULL_AUDIO_SAMPLES / chunk_size;
let chunk = vec![0i16; chunk_size];
for _ in 0..full_chunks {
let ok = append_stream_samples(&mut audio, &mut full, &chunk);
assert!(ok, "should succeed while under cap");
}
// The buffer may now be at or just below MAX_FULL_AUDIO_SAMPLES (depending on
// whether MAX_FULL_AUDIO_SAMPLES is an exact multiple of chunk_size).
assert!(
full.len() <= MAX_FULL_AUDIO_SAMPLES,
"full_audio_buf must not exceed cap before overflow chunk"
);
// One more chunk must be rejected.
let extra = vec![1i16; chunk_size];
let ok = append_stream_samples(&mut audio, &mut full, &extra);
assert!(
!ok,
"must return false (cap exceeded) when appending would overflow"
);
// The buffer must not have grown.
assert!(
full.len() <= MAX_FULL_AUDIO_SAMPLES,
"full_audio_buf must not exceed MAX_FULL_AUDIO_SAMPLES after cap is hit"
);
}
/// A single oversized chunk that would exceed the cap on its own must also be rejected.
#[test]
fn append_stream_samples_rejects_single_oversized_chunk() {
let mut audio = Vec::new();
let mut full = Vec::new();
// Pre-fill to near the cap (1 sample short).
let near_full = vec![0i16; MAX_FULL_AUDIO_SAMPLES - 1];
let ok = append_stream_samples(&mut audio, &mut full, &near_full);
assert!(ok, "pre-fill should succeed");
// A 2-sample chunk would push us 1 sample over the cap.
let ok = append_stream_samples(&mut audio, &mut full, &[7, 8]);
assert!(!ok, "must return false when chunk crosses the cap boundary");
assert!(
full.len() <= MAX_FULL_AUDIO_SAMPLES,
"full_audio_buf must not exceed cap"
);
}
#[test]
fn is_stop_command_only_accepts_stop_type() {
assert!(is_stop_command(r#"{"type":"stop"}"#));