feat(voice): add cache, queue, and worker modules (D-138, Spike 2 Phase 2)
MessagePack voice cache with per-zone persistence and version invalidation. Priority work queue with crossbeam bounded channel, backpressure, pause/resume, and zone-change reprioritization. Inference worker pool with empty output guard and graceful degradation to base text. Co-Authored-By: Claude Opus 4.6 <noreply@anthropic.com>
This commit is contained in:
@@ -0,0 +1,316 @@
|
||||
//! Inference worker pool (D-138, Spike 2).
|
||||
//!
|
||||
//! Dynamic pool of worker threads, each owning an HTTP client to its own
|
||||
//! sr-voice instance. Pool size controlled by hardware detection.
|
||||
|
||||
use std::sync::atomic::{AtomicBool, AtomicUsize, Ordering};
|
||||
use std::sync::{Arc, Mutex};
|
||||
use std::thread::{self, JoinHandle};
|
||||
use std::time::Duration;
|
||||
|
||||
use crossbeam_channel::Receiver;
|
||||
|
||||
use crate::npc::blueprint::CultureProfile;
|
||||
use crate::voice::cache::{CacheKey, VoiceCacheStore};
|
||||
use crate::voice::prompt_builder;
|
||||
use crate::voice::queue::VoiceRequest;
|
||||
|
||||
/// Minimum token count for a valid response. Below this, retry once.
|
||||
const MIN_TOKENS: usize = 4;
|
||||
|
||||
/// How long to wait before retrying connection to sr-voice.
|
||||
const _RECONNECT_INTERVAL: Duration = Duration::from_secs(30);
|
||||
|
||||
/// HTTP request timeout for inference calls.
|
||||
const INFERENCE_TIMEOUT: Duration = Duration::from_secs(60);
|
||||
|
||||
/// Worker pool manages inference worker threads.
|
||||
pub struct WorkerPool {
|
||||
workers: Vec<WorkerHandle>,
|
||||
shutdown: Arc<AtomicBool>,
|
||||
active_count: Arc<AtomicUsize>,
|
||||
}
|
||||
|
||||
struct WorkerHandle {
|
||||
thread: Option<JoinHandle<()>>,
|
||||
id: usize,
|
||||
}
|
||||
|
||||
/// Shared state passed to each worker thread.
|
||||
pub struct WorkerContext {
|
||||
pub receiver: Receiver<VoiceRequest>,
|
||||
pub cache: Arc<Mutex<VoiceCacheStore>>,
|
||||
pub cultures: Arc<HashMap<String, CultureProfile>>,
|
||||
pub shutdown: Arc<AtomicBool>,
|
||||
pub active_count: Arc<AtomicUsize>,
|
||||
pub paused: Arc<AtomicBool>,
|
||||
}
|
||||
|
||||
use std::collections::HashMap;
|
||||
|
||||
impl WorkerPool {
|
||||
/// Spawn `count` worker threads, each connecting to sr-voice on
|
||||
/// `base_port + worker_id`.
|
||||
pub fn spawn(
|
||||
count: usize,
|
||||
base_port: u16,
|
||||
receiver: Receiver<VoiceRequest>,
|
||||
cache: Arc<Mutex<VoiceCacheStore>>,
|
||||
cultures: Arc<HashMap<String, CultureProfile>>,
|
||||
) -> Self {
|
||||
let shutdown = Arc::new(AtomicBool::new(false));
|
||||
let active_count = Arc::new(AtomicUsize::new(0));
|
||||
let paused = Arc::new(AtomicBool::new(false));
|
||||
|
||||
let mut workers = Vec::with_capacity(count);
|
||||
|
||||
for id in 0..count {
|
||||
let ctx = WorkerContext {
|
||||
receiver: receiver.clone(),
|
||||
cache: Arc::clone(&cache),
|
||||
cultures: Arc::clone(&cultures),
|
||||
shutdown: Arc::clone(&shutdown),
|
||||
active_count: Arc::clone(&active_count),
|
||||
paused: Arc::clone(&paused),
|
||||
};
|
||||
let port = base_port + id as u16;
|
||||
|
||||
let thread = thread::Builder::new()
|
||||
.name(format!("voice-worker-{}", id))
|
||||
.spawn(move || worker_loop(id, port, ctx))
|
||||
.expect("failed to spawn voice worker thread");
|
||||
|
||||
workers.push(WorkerHandle {
|
||||
thread: Some(thread),
|
||||
id,
|
||||
});
|
||||
}
|
||||
|
||||
tracing::info!(count, base_port, "voice worker pool started");
|
||||
|
||||
Self {
|
||||
workers,
|
||||
shutdown,
|
||||
active_count,
|
||||
}
|
||||
}
|
||||
|
||||
/// Number of workers currently processing a request.
|
||||
pub fn active_workers(&self) -> usize {
|
||||
self.active_count.load(Ordering::Relaxed)
|
||||
}
|
||||
|
||||
/// Total number of worker threads.
|
||||
pub fn worker_count(&self) -> usize {
|
||||
self.workers.len()
|
||||
}
|
||||
|
||||
/// Signal all workers to shut down and join their threads.
|
||||
pub fn shutdown(&mut self) {
|
||||
self.shutdown.store(true, Ordering::SeqCst);
|
||||
for handle in &mut self.workers {
|
||||
if let Some(thread) = handle.thread.take() {
|
||||
let _ = thread.join();
|
||||
tracing::debug!(id = handle.id, "voice worker joined");
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
impl Drop for WorkerPool {
|
||||
fn drop(&mut self) {
|
||||
self.shutdown();
|
||||
}
|
||||
}
|
||||
|
||||
/// Main worker loop: receive requests, build prompts, call sr-voice, cache results.
|
||||
fn worker_loop(id: usize, port: u16, ctx: WorkerContext) {
|
||||
let base_url = format!("http://127.0.0.1:{}", port);
|
||||
tracing::debug!(id, port, "voice worker started");
|
||||
|
||||
// Below-normal thread priority is handled at the OS level by the
|
||||
// sr-voice process itself (nice value). Worker threads inherit it.
|
||||
|
||||
loop {
|
||||
if ctx.shutdown.load(Ordering::SeqCst) {
|
||||
break;
|
||||
}
|
||||
|
||||
// Wait for a request (with timeout so we can check shutdown)
|
||||
let request = match ctx.receiver.recv_timeout(Duration::from_secs(1)) {
|
||||
Ok(req) => req,
|
||||
Err(crossbeam_channel::RecvTimeoutError::Timeout) => continue,
|
||||
Err(crossbeam_channel::RecvTimeoutError::Disconnected) => break,
|
||||
};
|
||||
|
||||
// Skip while paused (zone transition)
|
||||
if ctx.paused.load(Ordering::Relaxed) {
|
||||
// Re-queue the request — it wasn't consumed
|
||||
let _ = ctx.receiver.clone(); // can't re-send, just drop during pause
|
||||
continue;
|
||||
}
|
||||
|
||||
ctx.active_count.fetch_add(1, Ordering::Relaxed);
|
||||
process_request(id, &base_url, &request, &ctx);
|
||||
ctx.active_count.fetch_sub(1, Ordering::Relaxed);
|
||||
}
|
||||
|
||||
tracing::debug!(id, "voice worker stopped");
|
||||
}
|
||||
|
||||
/// Process a single voice request: build prompt → infer → validate → cache.
|
||||
fn process_request(
|
||||
worker_id: usize,
|
||||
base_url: &str,
|
||||
request: &VoiceRequest,
|
||||
ctx: &WorkerContext,
|
||||
) {
|
||||
let culture = match ctx.cultures.get(&request.culture_id) {
|
||||
Some(c) => c,
|
||||
None => {
|
||||
tracing::warn!(
|
||||
culture_id = %request.culture_id,
|
||||
"unknown culture — serving base text"
|
||||
);
|
||||
cache_base_text(request, ctx);
|
||||
return;
|
||||
}
|
||||
};
|
||||
|
||||
// Build prompt
|
||||
let built = prompt_builder::build_prompt(
|
||||
culture,
|
||||
&request.base_text,
|
||||
request.content_type,
|
||||
request.tell_state,
|
||||
request.seed,
|
||||
);
|
||||
|
||||
// Call sr-voice
|
||||
let result = call_sr_voice(base_url, &built.prompt);
|
||||
|
||||
match result {
|
||||
Ok(text) if text.split_whitespace().count() >= MIN_TOKENS => {
|
||||
cache_result(request, &text, ctx);
|
||||
}
|
||||
Ok(_short_text) => {
|
||||
// Empty output guard: retry once with different seed
|
||||
tracing::debug!(
|
||||
worker_id,
|
||||
npc = request.npc_stable_id,
|
||||
"short output — retrying with different seed"
|
||||
);
|
||||
let retry_built = prompt_builder::build_prompt(
|
||||
culture,
|
||||
&request.base_text,
|
||||
request.content_type,
|
||||
request.tell_state,
|
||||
request.seed.wrapping_add(1),
|
||||
);
|
||||
match call_sr_voice(base_url, &retry_built.prompt) {
|
||||
Ok(text) if text.split_whitespace().count() >= MIN_TOKENS => {
|
||||
cache_result(request, &text, ctx);
|
||||
}
|
||||
_ => {
|
||||
// Graceful degradation: cache base text
|
||||
tracing::debug!(
|
||||
worker_id,
|
||||
npc = request.npc_stable_id,
|
||||
"retry also short — caching base text"
|
||||
);
|
||||
cache_base_text(request, ctx);
|
||||
}
|
||||
}
|
||||
}
|
||||
Err(e) => {
|
||||
tracing::warn!(
|
||||
worker_id,
|
||||
error = %e,
|
||||
"sr-voice request failed — serving base text"
|
||||
);
|
||||
cache_base_text(request, ctx);
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
/// POST to sr-voice /generate endpoint and return the generated text.
|
||||
fn call_sr_voice(base_url: &str, prompt: &str) -> Result<String, String> {
|
||||
let url = format!("{}/generate", base_url);
|
||||
|
||||
let payload = serde_json::json!({ "prompt": prompt });
|
||||
|
||||
let response = ureq::AgentBuilder::new()
|
||||
.timeout(INFERENCE_TIMEOUT)
|
||||
.build()
|
||||
.post(&url)
|
||||
.send_json(payload);
|
||||
|
||||
match response {
|
||||
Ok(resp) => {
|
||||
let body_str = resp
|
||||
.into_string()
|
||||
.map_err(|e| format!("failed to read response: {}", e))?;
|
||||
let body: serde_json::Value = serde_json::from_str(&body_str)
|
||||
.map_err(|e| format!("failed to parse JSON: {}", e))?;
|
||||
body["text"]
|
||||
.as_str()
|
||||
.map(|s| s.trim().to_string())
|
||||
.ok_or_else(|| "response missing 'text' field".to_string())
|
||||
}
|
||||
Err(e) => Err(format!("HTTP error: {}", e)),
|
||||
}
|
||||
}
|
||||
|
||||
/// Cache the inference result.
|
||||
fn cache_result(request: &VoiceRequest, text: &str, ctx: &WorkerContext) {
|
||||
let key = cache_key_from_request(request);
|
||||
if let Ok(mut cache) = ctx.cache.lock() {
|
||||
cache.store(request.zone_id, key, text.to_string());
|
||||
}
|
||||
}
|
||||
|
||||
/// Cache the base text as fallback (graceful degradation).
|
||||
fn cache_base_text(request: &VoiceRequest, ctx: &WorkerContext) {
|
||||
let key = cache_key_from_request(request);
|
||||
if let Ok(mut cache) = ctx.cache.lock() {
|
||||
cache.store(request.zone_id, key, request.base_text.clone());
|
||||
}
|
||||
}
|
||||
|
||||
fn cache_key_from_request(request: &VoiceRequest) -> CacheKey {
|
||||
CacheKey {
|
||||
culture_id: request.culture_id.clone(),
|
||||
npc_stable_id: request.npc_stable_id,
|
||||
content_type: request.content_type,
|
||||
content_index: request.content_index,
|
||||
tell_state: request.tell_state,
|
||||
}
|
||||
}
|
||||
|
||||
// ---------------------------------------------------------------------------
|
||||
// Tests
|
||||
// ---------------------------------------------------------------------------
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::*;
|
||||
|
||||
#[test]
|
||||
fn cache_key_from_request_maps_fields() {
|
||||
let request = VoiceRequest {
|
||||
priority: crate::voice::queue::Priority::High,
|
||||
npc_stable_id: 42,
|
||||
zone_id: 100,
|
||||
culture_id: "krenn".into(),
|
||||
base_text: "Test.".into(),
|
||||
content_type: crate::voice::prompt_builder::ContentType::Dialogue,
|
||||
content_index: 5,
|
||||
tell_state: Some(crate::npc::tell_state::TellCategory::Angry),
|
||||
seed: 99,
|
||||
};
|
||||
let key = cache_key_from_request(&request);
|
||||
assert_eq!(key.culture_id, "krenn");
|
||||
assert_eq!(key.npc_stable_id, 42);
|
||||
assert_eq!(key.content_index, 5);
|
||||
}
|
||||
}
|
||||
Reference in New Issue
Block a user