#![allow(dead_code)]
use log::{debug, info, warn};
use std::sync::Arc;
use std::sync::atomic::{AtomicBool, AtomicUsize, Ordering};
use std::time::Duration;
use tokio::sync::{Mutex, RwLock, mpsc};
use super::config::WorkerConfig;
use super::queue::PriorityRequestQueue;
use super::response_channel::ResponseChannel;
use crate::domain::EmbedResponse;
use crate::error::VecboostError;
use crate::service::embedding::EmbeddingService;
#[derive(Debug)]
pub enum WorkerTask {
ProcessRequest {
request_id: String,
embed_request: crate::domain::EmbedRequest,
},
Shutdown {
immediate: bool,
},
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum WorkerState {
Idle,
Processing,
Stopping,
Stopped,
}
pub struct Worker {
worker_id: usize,
running: Arc<AtomicBool>,
state: Arc<Mutex<WorkerState>>,
receiver: mpsc::Receiver<WorkerTask>,
config: WorkerConfig,
}
pub struct WorkerManager {
min_workers: usize,
max_workers: usize,
current_workers: Arc<AtomicUsize>,
queue: Arc<PriorityRequestQueue>,
response_channel: Arc<ResponseChannel>,
config: WorkerConfig,
running: Arc<AtomicBool>,
worker_senders: Arc<Mutex<Vec<mpsc::Sender<WorkerTask>>>>,
embedding_service: Arc<RwLock<EmbeddingService>>,
worker_health: Arc<Mutex<Vec<WorkerHealthInfo>>>,
}
#[derive(Debug, Clone)]
struct WorkerHealthInfo {
worker_id: usize,
last_active_time: std::time::Instant,
crash_count: usize,
is_alive: bool,
}
impl WorkerHealthInfo {
fn new(worker_id: usize) -> Self {
Self {
worker_id,
last_active_time: std::time::Instant::now(),
crash_count: 0,
is_alive: true,
}
}
fn update_activity(&mut self) {
self.last_active_time = std::time::Instant::now();
}
fn record_crash(&mut self) {
self.crash_count += 1;
}
}
impl WorkerManager {
pub fn new(
queue: Arc<PriorityRequestQueue>,
response_channel: Arc<ResponseChannel>,
config: WorkerConfig,
embedding_service: Arc<RwLock<EmbeddingService>>,
) -> Self {
Self {
min_workers: config.min_workers,
max_workers: config.max_workers,
current_workers: Arc::new(AtomicUsize::new(0)),
queue,
response_channel,
config,
running: Arc::new(AtomicBool::new(true)),
worker_senders: Arc::new(Mutex::new(Vec::new())),
embedding_service,
worker_health: Arc::new(Mutex::new(Vec::new())),
}
}
pub async fn start(&self) -> Result<(), VecboostError> {
info!(
"Starting WorkerManager with min={} max={}",
self.min_workers, self.max_workers
);
for _ in 0..self.min_workers {
self.spawn_worker().await;
}
self.start_scaling_monitor().await;
info!("WorkerManager started successfully");
Ok(())
}
pub async fn shutdown(&self) {
info!("Shutting down WorkerManager...");
self.running.store(false, Ordering::SeqCst);
let senders = {
let guard = self.worker_senders.lock().await;
guard.clone()
};
for sender in &senders {
let _ = sender.send(WorkerTask::Shutdown { immediate: false }).await;
}
tokio::time::sleep(Duration::from_secs(5)).await;
for sender in &senders {
let _ = sender.send(WorkerTask::Shutdown { immediate: true }).await;
}
info!("WorkerManager shutdown complete");
}
pub fn current_workers(&self) -> usize {
self.current_workers.load(Ordering::SeqCst)
}
pub async fn spawn_worker(&self) {
Self::spawn_single_worker(
&self.current_workers,
&self.worker_senders,
&self.worker_health,
&self.queue,
&self.response_channel,
&self.embedding_service,
&self.config,
&self.running,
)
.await;
}
#[allow(clippy::too_many_arguments)]
async fn spawn_single_worker(
current_workers: &Arc<AtomicUsize>,
worker_senders: &Arc<Mutex<Vec<mpsc::Sender<WorkerTask>>>>,
worker_health: &Arc<Mutex<Vec<WorkerHealthInfo>>>,
queue: &Arc<PriorityRequestQueue>,
response_channel: &Arc<ResponseChannel>,
embedding_service: &Arc<RwLock<EmbeddingService>>,
config: &WorkerConfig,
running: &Arc<AtomicBool>,
) {
let worker_id = current_workers.fetch_add(1, Ordering::SeqCst);
let (task_sender, task_receiver) = mpsc::channel(100);
{
let mut senders = worker_senders.lock().await;
senders.push(task_sender.clone());
}
{
let mut health_guard = worker_health.lock().await;
health_guard.push(WorkerHealthInfo::new(worker_id));
}
let queue = Arc::clone(queue);
let response_channel = Arc::clone(response_channel);
let config = config.clone();
let running = Arc::clone(running);
let embedding_service = Arc::clone(embedding_service);
let worker_health = Arc::clone(worker_health);
let current_workers = Arc::clone(current_workers);
info!("Worker {} started", worker_id);
tokio::spawn(async move {
Self::worker_loop(
worker_id,
task_receiver,
queue,
response_channel,
config,
running,
embedding_service,
worker_health,
current_workers,
)
.await;
});
}
#[allow(clippy::too_many_arguments)]
async fn worker_loop(
worker_id: usize,
mut task_receiver: mpsc::Receiver<WorkerTask>,
queue: Arc<PriorityRequestQueue>,
response_channel: Arc<ResponseChannel>,
config: WorkerConfig,
running: Arc<AtomicBool>,
embedding_service: Arc<RwLock<EmbeddingService>>,
worker_health: Arc<Mutex<Vec<WorkerHealthInfo>>>,
current_workers: Arc<AtomicUsize>,
) {
debug!("Worker {} loop started", worker_id);
let mut idle_count: usize = 0;
const MAX_IDLE_COUNT: usize = 10;
loop {
if !running.load(Ordering::Relaxed) {
info!("Worker {} received stop signal", worker_id);
break;
}
{
let mut guard = worker_health.lock().await;
if let Some(info) = guard.iter_mut().find(|i| i.worker_id == worker_id) {
info.update_activity();
}
}
tokio::select! {
Some(task) = task_receiver.recv() => {
match task {
WorkerTask::Shutdown { immediate } => {
if immediate {
info!(
"Worker {} received immediate shutdown signal",
worker_id
);
} else {
info!(
"Worker {} received graceful shutdown signal",
worker_id
);
}
break;
}
WorkerTask::ProcessRequest { .. } => {
debug!(
"Worker {} received ProcessRequest via task_receiver \
(ignored, queue is primary route)",
worker_id
);
}
}
}
Some(request) = queue.dequeue() => {
idle_count = 0;
debug!(
"Worker {} processing request {}",
worker_id, request.request_id
);
let result = Self::process_request(&request, &embedding_service).await;
response_channel
.complete(request.request_id.clone(), result)
.await;
debug!(
"Worker {} completed request {}",
worker_id, request.request_id
);
}
else => {
idle_count = idle_count.saturating_add(1usize);
let wait_ms =
std::cmp::min(100usize * (1usize << idle_count.min(6)), 5000usize);
debug!(
"Worker {} queue empty, waiting {}ms (idle_count={})",
worker_id, wait_ms, idle_count
);
tokio::time::sleep(Duration::from_millis(wait_ms as u64)).await;
if idle_count > MAX_IDLE_COUNT && worker_id > config.min_workers {
info!(
"Worker {} idle for too long, requesting shutdown",
worker_id
);
break;
}
}
}
}
let final_count = Self::decrement_worker_count(¤t_workers);
info!(
"Worker {} stopped, remaining workers: {}",
worker_id, final_count
);
{
let mut guard = worker_health.lock().await;
if let Some(info) = guard.iter_mut().find(|i| i.worker_id == worker_id) {
info.is_alive = false;
}
}
}
fn decrement_worker_count(current_workers: &Arc<AtomicUsize>) -> usize {
current_workers.fetch_sub(1, Ordering::SeqCst) - 1
}
async fn process_request(
request: &super::queue::QueuedRequest,
embedding_service: &Arc<RwLock<EmbeddingService>>,
) -> Result<EmbedResponse, VecboostError> {
let embed_request = &request.embed_request;
debug!("Processing embedding request");
let service_guard = embedding_service.read().await;
let result = service_guard
.process_text(
crate::domain::EmbedRequest {
text: embed_request.text.clone(),
normalize: embed_request.normalize,
},
None, )
.await;
drop(service_guard);
match result {
Ok(response) => {
debug!(
"Successfully generated embedding with dimension: {}",
response.dimension
);
Ok(response)
}
Err(e) => {
warn!("Embedding inference failed: {}", e);
Err(e)
}
}
}
async fn start_scaling_monitor(&self) {
let queue = Arc::clone(&self.queue);
let current_workers = Arc::clone(&self.current_workers);
let config = self.config.clone();
let running = Arc::clone(&self.running);
let worker_senders = Arc::clone(&self.worker_senders);
let worker_health = Arc::clone(&self.worker_health);
let response_channel = Arc::clone(&self.response_channel);
let embedding_service = Arc::clone(&self.embedding_service);
tokio::spawn(async move {
let mut interval =
tokio::time::interval(Duration::from_secs(config.scale_check_interval_secs));
loop {
if !running.load(Ordering::Relaxed) {
break;
}
interval.tick().await;
let queue_size = queue.size();
let current = current_workers.load(Ordering::SeqCst);
if queue_size > config.scale_up_threshold && current < config.max_workers {
let new_workers = std::cmp::min(
(queue_size / config.scale_up_threshold).saturating_sub(1),
config.max_workers - current,
);
if new_workers > 0 {
info!(
"Scaling up: adding {} workers (queue size: {})",
new_workers, queue_size
);
for _ in 0..new_workers {
Self::spawn_single_worker(
¤t_workers,
&worker_senders,
&worker_health,
&queue,
&response_channel,
&embedding_service,
&config,
&running,
)
.await;
}
}
}
if queue_size < config.scale_down_threshold && current > config.min_workers {
let mut senders = worker_senders.lock().await;
let before = senders.len();
senders.retain(|s| !s.is_closed());
let cleaned = before - senders.len();
if cleaned > 0 {
debug!(
"Cleaned {} stale worker senders (before: {}, after: {})",
cleaned,
before,
senders.len()
);
}
let to_remove = current - config.min_workers;
if to_remove == 0 || senders.is_empty() {
continue;
}
info!(
"Scaling down: removing {} workers (queue size: {})",
to_remove, queue_size
);
for sender in senders.iter().rev().take(to_remove) {
let _ = sender.send(WorkerTask::Shutdown { immediate: false }).await;
}
}
}
});
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::config::model::{ModelConfig, Precision};
use crate::domain::EmbedRequest;
use crate::engine::InferenceEngine;
use crate::pipeline::priority::{Priority, RequestSource};
use crate::pipeline::queue::QueuedRequest;
use async_trait::async_trait;
struct MockEngine;
#[async_trait]
impl InferenceEngine for MockEngine {
fn embed(&self, _text: &str) -> Result<Vec<f32>, VecboostError> {
Ok(vec![0.0; 8])
}
fn embed_batch(&self, texts: &[String]) -> Result<Vec<Vec<f32>>, VecboostError> {
Ok(texts.iter().map(|_| vec![0.0; 8]).collect())
}
fn precision(&self) -> &Precision {
&Precision::Fp32
}
fn supports_mixed_precision(&self) -> bool {
false
}
async fn try_fallback_to_cpu(
&mut self,
_config: &ModelConfig,
) -> Result<(), VecboostError> {
Ok(())
}
}
struct ErrorEngine;
#[async_trait]
impl InferenceEngine for ErrorEngine {
fn embed(&self, _text: &str) -> Result<Vec<f32>, VecboostError> {
Err(VecboostError::InferenceError(
"mock inference failure".to_string(),
))
}
fn embed_batch(&self, _texts: &[String]) -> Result<Vec<Vec<f32>>, VecboostError> {
Err(VecboostError::InferenceError(
"mock batch inference failure".to_string(),
))
}
fn precision(&self) -> &Precision {
&Precision::Fp32
}
fn supports_mixed_precision(&self) -> bool {
false
}
async fn try_fallback_to_cpu(
&mut self,
_config: &ModelConfig,
) -> Result<(), VecboostError> {
Ok(())
}
}
#[test]
fn test_decrement_worker_count_actually_decrements() {
let counter = Arc::new(AtomicUsize::new(5));
let remaining = WorkerManager::decrement_worker_count(&counter);
assert_eq!(
remaining, 4,
"decrement_worker_count must return new count (old-1), not hardcoded 0"
);
assert_eq!(
counter.load(Ordering::SeqCst),
4,
"counter must be decremented from 5 to 4"
);
}
#[test]
fn test_decrement_worker_count_from_one_to_zero() {
let counter = Arc::new(AtomicUsize::new(1));
let remaining = WorkerManager::decrement_worker_count(&counter);
assert_eq!(remaining, 0, "single worker stop should bring count to 0");
assert_eq!(counter.load(Ordering::SeqCst), 0);
}
#[test]
fn test_decrement_worker_count_multiple_times() {
let counter = Arc::new(AtomicUsize::new(3));
assert_eq!(WorkerManager::decrement_worker_count(&counter), 2);
assert_eq!(WorkerManager::decrement_worker_count(&counter), 1);
assert_eq!(WorkerManager::decrement_worker_count(&counter), 0);
assert_eq!(counter.load(Ordering::SeqCst), 0);
}
#[tokio::test]
async fn test_spawn_worker_increments_current_workers() {
let queue = Arc::new(PriorityRequestQueue::new(100));
let response_channel = Arc::new(ResponseChannel::new());
let config = WorkerConfig::default();
let engine: Arc<RwLock<dyn InferenceEngine + Send + Sync>> =
Arc::new(RwLock::new(MockEngine));
let service = Arc::new(RwLock::new(EmbeddingService::new(engine, None)));
let manager = WorkerManager::new(queue, response_channel, config, service);
assert_eq!(manager.current_workers(), 0, "initial count must be 0");
manager.spawn_worker().await;
assert_eq!(
manager.current_workers(),
1,
"spawn_worker must increment counter via fetch_add (H2 fix)"
);
manager.spawn_worker().await;
assert_eq!(manager.current_workers(), 2);
manager.running.store(false, Ordering::SeqCst);
tokio::time::sleep(Duration::from_millis(50)).await;
}
#[tokio::test]
async fn test_spawn_single_worker_increments_counter() {
let queue = Arc::new(PriorityRequestQueue::new(100));
let response_channel = Arc::new(ResponseChannel::new());
let config = WorkerConfig::default();
let engine: Arc<RwLock<dyn InferenceEngine + Send + Sync>> =
Arc::new(RwLock::new(MockEngine));
let service = Arc::new(RwLock::new(EmbeddingService::new(engine, None)));
let manager = WorkerManager::new(queue, response_channel, config, service);
let running = Arc::clone(&manager.running);
let current_workers = Arc::clone(&manager.current_workers);
let worker_senders = Arc::clone(&manager.worker_senders);
let worker_health = Arc::clone(&manager.worker_health);
let queue_clone = Arc::clone(&manager.queue);
let response_channel_clone = Arc::clone(&manager.response_channel);
let embedding_service_clone = Arc::clone(&manager.embedding_service);
let config_clone = manager.config.clone();
assert_eq!(current_workers.load(Ordering::SeqCst), 0);
WorkerManager::spawn_single_worker(
¤t_workers,
&worker_senders,
&worker_health,
&queue_clone,
&response_channel_clone,
&embedding_service_clone,
&config_clone,
&running,
)
.await;
assert_eq!(
current_workers.load(Ordering::SeqCst),
1,
"spawn_single_worker (used by scaling monitor) must increment counter"
);
running.store(false, Ordering::SeqCst);
tokio::time::sleep(Duration::from_millis(50)).await;
}
#[tokio::test]
async fn test_worker_loop_consumes_graceful_shutdown_within_2s() {
let queue = Arc::new(PriorityRequestQueue::new(100));
let response_channel = Arc::new(ResponseChannel::new());
let config = WorkerConfig::default();
let engine: Arc<RwLock<dyn InferenceEngine + Send + Sync>> =
Arc::new(RwLock::new(MockEngine));
let service = Arc::new(RwLock::new(EmbeddingService::new(engine, None)));
let manager = WorkerManager::new(queue, response_channel, config, service);
manager.spawn_worker().await;
assert_eq!(manager.current_workers(), 1, "worker must be spawned");
{
let senders = manager.worker_senders.lock().await;
assert!(
!senders.is_empty(),
"worker_senders must contain spawned worker's sender"
);
senders[0]
.send(WorkerTask::Shutdown { immediate: false })
.await
.expect("send graceful Shutdown must succeed");
}
let deadline = tokio::time::Instant::now() + Duration::from_secs(2);
loop {
if manager.current_workers() == 0 {
break;
}
if tokio::time::Instant::now() >= deadline {
panic!(
"worker did not shut down within 2s after graceful Shutdown signal \
(current_workers={}) — worker_loop is not consuming task_receiver (T009 regression)",
manager.current_workers()
);
}
tokio::time::sleep(Duration::from_millis(50)).await;
}
}
#[tokio::test]
async fn test_worker_loop_consumes_immediate_shutdown_within_2s() {
let queue = Arc::new(PriorityRequestQueue::new(100));
let response_channel = Arc::new(ResponseChannel::new());
let config = WorkerConfig::default();
let engine: Arc<RwLock<dyn InferenceEngine + Send + Sync>> =
Arc::new(RwLock::new(MockEngine));
let service = Arc::new(RwLock::new(EmbeddingService::new(engine, None)));
let manager = WorkerManager::new(queue, response_channel, config, service);
manager.spawn_worker().await;
assert_eq!(manager.current_workers(), 1);
{
let senders = manager.worker_senders.lock().await;
senders[0]
.send(WorkerTask::Shutdown { immediate: true })
.await
.expect("send immediate Shutdown must succeed");
}
let deadline = tokio::time::Instant::now() + Duration::from_secs(2);
loop {
if manager.current_workers() == 0 {
break;
}
if tokio::time::Instant::now() >= deadline {
panic!(
"worker did not shut down within 2s after immediate Shutdown \
(current_workers={}) — worker_loop is not consuming task_receiver",
manager.current_workers()
);
}
tokio::time::sleep(Duration::from_millis(50)).await;
}
}
#[tokio::test]
async fn test_worker_exit_marks_sender_closed_for_cleanup() {
let queue = Arc::new(PriorityRequestQueue::new(100));
let response_channel = Arc::new(ResponseChannel::new());
let config = WorkerConfig::default();
let engine: Arc<RwLock<dyn InferenceEngine + Send + Sync>> =
Arc::new(RwLock::new(MockEngine));
let service = Arc::new(RwLock::new(EmbeddingService::new(engine, None)));
let manager = WorkerManager::new(queue, response_channel, config, service);
manager.spawn_worker().await;
assert_eq!(manager.current_workers(), 1);
{
let senders = manager.worker_senders.lock().await;
senders[0]
.send(WorkerTask::Shutdown { immediate: false })
.await
.expect("send Shutdown must succeed");
}
let deadline = tokio::time::Instant::now() + Duration::from_secs(2);
loop {
if manager.current_workers() == 0 {
break;
}
if tokio::time::Instant::now() >= deadline {
panic!("worker did not exit within 2s");
}
tokio::time::sleep(Duration::from_millis(50)).await;
}
{
let senders = manager.worker_senders.lock().await;
assert!(
senders[0].is_closed(),
"sender must be closed after worker exit (required for retain cleanup in scaling_monitor)"
);
}
}
#[tokio::test]
async fn test_worker_manager_new_initial_state() {
let queue = Arc::new(PriorityRequestQueue::new(100));
let response_channel = Arc::new(ResponseChannel::new());
let config = WorkerConfig::default();
let engine: Arc<RwLock<dyn InferenceEngine + Send + Sync>> =
Arc::new(RwLock::new(MockEngine));
let service = Arc::new(RwLock::new(EmbeddingService::new(engine, None)));
let manager = WorkerManager::new(queue, response_channel, config, service);
assert_eq!(manager.current_workers(), 0);
assert!(manager.running.load(Ordering::SeqCst));
assert!(manager.worker_senders.lock().await.is_empty());
assert!(manager.worker_health.lock().await.is_empty());
}
#[tokio::test(flavor = "multi_thread")]
async fn test_worker_manager_start_spawns_min_workers() {
let queue = Arc::new(PriorityRequestQueue::new(100));
let response_channel = Arc::new(ResponseChannel::new());
let config = WorkerConfig {
min_workers: 3,
max_workers: 8,
..Default::default()
};
let engine: Arc<RwLock<dyn InferenceEngine + Send + Sync>> =
Arc::new(RwLock::new(MockEngine));
let service = Arc::new(RwLock::new(EmbeddingService::new(engine, None)));
let manager = WorkerManager::new(queue, response_channel, config, service);
manager.start().await.unwrap();
assert_eq!(
manager.current_workers(),
3,
"start() must spawn min_workers workers"
);
assert_eq!(
manager.worker_senders.lock().await.len(),
3,
"must have 3 senders"
);
assert_eq!(
manager.worker_health.lock().await.len(),
3,
"must have 3 health entries"
);
manager.running.store(false, Ordering::SeqCst);
let senders = manager.worker_senders.lock().await.clone();
for s in &senders {
let _ = s.send(WorkerTask::Shutdown { immediate: true }).await;
}
tokio::time::sleep(Duration::from_millis(100)).await;
}
#[tokio::test(flavor = "multi_thread")]
async fn test_worker_manager_shutdown_stops_workers() {
let queue = Arc::new(PriorityRequestQueue::new(100));
let response_channel = Arc::new(ResponseChannel::new());
let config = WorkerConfig::default();
let engine: Arc<RwLock<dyn InferenceEngine + Send + Sync>> =
Arc::new(RwLock::new(MockEngine));
let service = Arc::new(RwLock::new(EmbeddingService::new(engine, None)));
let manager = WorkerManager::new(queue, response_channel, config, service);
manager.spawn_worker().await;
manager.spawn_worker().await;
assert_eq!(manager.current_workers(), 2);
tokio::time::timeout(Duration::from_secs(15), manager.shutdown())
.await
.expect("shutdown should complete within 15s");
assert!(
!manager.running.load(Ordering::SeqCst),
"running flag must be false after shutdown"
);
let deadline = tokio::time::Instant::now() + Duration::from_secs(3);
loop {
if manager.current_workers() == 0 {
break;
}
if tokio::time::Instant::now() >= deadline {
panic!(
"workers did not exit after shutdown, remaining: {}",
manager.current_workers()
);
}
tokio::time::sleep(Duration::from_millis(50)).await;
}
}
#[tokio::test]
async fn test_process_request_success() {
let engine: Arc<RwLock<dyn InferenceEngine + Send + Sync>> =
Arc::new(RwLock::new(MockEngine));
let service = Arc::new(RwLock::new(EmbeddingService::new(engine, None)));
let (tx, _rx) = tokio::sync::oneshot::channel();
let request = QueuedRequest {
request_id: "test-process-1".to_string(),
embed_request: EmbedRequest {
text: "hello world".to_string(),
normalize: Some(true),
},
priority: Priority::Normal,
submitted_at: std::time::Instant::now(),
timeout: Duration::from_secs(30),
source: RequestSource::http("127.0.0.1".to_string()),
response_tx: tx,
};
let result = WorkerManager::process_request(&request, &service).await;
assert!(result.is_ok(), "process_request should succeed");
let response = result.unwrap();
assert_eq!(response.dimension, 8);
assert_eq!(response.embedding.len(), 8);
}
#[tokio::test]
async fn test_process_request_engine_error() {
let engine: Arc<RwLock<dyn InferenceEngine + Send + Sync>> =
Arc::new(RwLock::new(ErrorEngine));
let service = Arc::new(RwLock::new(EmbeddingService::new(engine, None)));
let (tx, _rx) = tokio::sync::oneshot::channel();
let request = QueuedRequest {
request_id: "test-process-err".to_string(),
embed_request: EmbedRequest {
text: "hello".to_string(),
normalize: Some(true),
},
priority: Priority::Normal,
submitted_at: std::time::Instant::now(),
timeout: Duration::from_secs(30),
source: RequestSource::http("127.0.0.1".to_string()),
response_tx: tx,
};
let result = WorkerManager::process_request(&request, &service).await;
assert!(result.is_err());
match result.unwrap_err() {
VecboostError::InferenceError(msg) => {
assert!(msg.contains("mock inference failure"));
}
other => panic!("expected InferenceError, got {:?}", other),
}
}
#[tokio::test(flavor = "multi_thread")]
async fn test_worker_loop_processes_queued_request() {
let queue = Arc::new(PriorityRequestQueue::new(100));
let response_channel = Arc::new(ResponseChannel::new());
let config = WorkerConfig::default();
let engine: Arc<RwLock<dyn InferenceEngine + Send + Sync>> =
Arc::new(RwLock::new(MockEngine));
let service = Arc::new(RwLock::new(EmbeddingService::new(engine, None)));
let manager = WorkerManager::new(
Arc::clone(&queue),
Arc::clone(&response_channel),
config,
service,
);
let rx = response_channel.register("test-loop-1".to_string()).await;
let (tx, _) = tokio::sync::oneshot::channel();
let request = QueuedRequest {
request_id: "test-loop-1".to_string(),
embed_request: EmbedRequest {
text: "hello world".to_string(),
normalize: Some(true),
},
priority: Priority::Normal,
submitted_at: std::time::Instant::now(),
timeout: Duration::from_secs(30),
source: RequestSource::http("127.0.0.1".to_string()),
response_tx: tx,
};
queue.enqueue(request).await.unwrap();
manager.spawn_worker().await;
let result = tokio::time::timeout(Duration::from_secs(5), rx).await;
assert!(result.is_ok(), "response should arrive within 5s");
let response_result = result.unwrap().unwrap();
assert!(response_result.is_ok());
let response = response_result.unwrap();
assert_eq!(response.dimension, 8);
manager.running.store(false, Ordering::SeqCst);
let senders = manager.worker_senders.lock().await.clone();
for s in &senders {
let _ = s.send(WorkerTask::Shutdown { immediate: true }).await;
}
tokio::time::sleep(Duration::from_millis(100)).await;
}
#[tokio::test(flavor = "multi_thread")]
async fn test_worker_loop_propagates_engine_error() {
let queue = Arc::new(PriorityRequestQueue::new(100));
let response_channel = Arc::new(ResponseChannel::new());
let config = WorkerConfig::default();
let engine: Arc<RwLock<dyn InferenceEngine + Send + Sync>> =
Arc::new(RwLock::new(ErrorEngine));
let service = Arc::new(RwLock::new(EmbeddingService::new(engine, None)));
let manager = WorkerManager::new(
Arc::clone(&queue),
Arc::clone(&response_channel),
config,
service,
);
let rx = response_channel.register("test-loop-err".to_string()).await;
let (tx, _) = tokio::sync::oneshot::channel();
let request = QueuedRequest {
request_id: "test-loop-err".to_string(),
embed_request: EmbedRequest {
text: "hello".to_string(),
normalize: Some(true),
},
priority: Priority::Normal,
submitted_at: std::time::Instant::now(),
timeout: Duration::from_secs(30),
source: RequestSource::http("127.0.0.1".to_string()),
response_tx: tx,
};
queue.enqueue(request).await.unwrap();
manager.spawn_worker().await;
let result = tokio::time::timeout(Duration::from_secs(5), rx).await;
assert!(result.is_ok(), "response should arrive within 5s");
let response_result = result.unwrap().unwrap();
assert!(response_result.is_err());
match response_result.unwrap_err() {
VecboostError::InferenceError(msg) => {
assert!(msg.contains("mock inference failure"));
}
other => panic!("expected InferenceError, got {:?}", other),
}
manager.running.store(false, Ordering::SeqCst);
let senders = manager.worker_senders.lock().await.clone();
for s in &senders {
let _ = s.send(WorkerTask::Shutdown { immediate: true }).await;
}
tokio::time::sleep(Duration::from_millis(100)).await;
}
#[tokio::test(flavor = "multi_thread")]
async fn test_worker_exit_marks_health_dead() {
let queue = Arc::new(PriorityRequestQueue::new(100));
let response_channel = Arc::new(ResponseChannel::new());
let config = WorkerConfig::default();
let engine: Arc<RwLock<dyn InferenceEngine + Send + Sync>> =
Arc::new(RwLock::new(MockEngine));
let service = Arc::new(RwLock::new(EmbeddingService::new(engine, None)));
let manager = WorkerManager::new(queue, response_channel, config, service);
manager.spawn_worker().await;
assert_eq!(manager.current_workers(), 1);
{
let senders = manager.worker_senders.lock().await;
senders[0]
.send(WorkerTask::Shutdown { immediate: false })
.await
.unwrap();
}
let deadline = tokio::time::Instant::now() + Duration::from_secs(2);
loop {
if manager.current_workers() == 0 {
break;
}
if tokio::time::Instant::now() >= deadline {
panic!("worker did not exit within 2s");
}
tokio::time::sleep(Duration::from_millis(50)).await;
}
let health = manager.worker_health.lock().await;
assert_eq!(health.len(), 1);
assert!(
!health[0].is_alive,
"worker must be marked as dead after exit"
);
}
#[test]
fn test_worker_health_info_new() {
let health = WorkerHealthInfo::new(42);
assert_eq!(health.worker_id, 42);
assert_eq!(health.crash_count, 0);
assert!(health.is_alive);
}
#[test]
fn test_worker_health_info_update_activity() {
let mut health = WorkerHealthInfo::new(0);
let original = health.last_active_time;
std::thread::sleep(Duration::from_millis(5));
health.update_activity();
assert!(health.last_active_time > original);
}
#[test]
fn test_worker_health_info_record_crash() {
let mut health = WorkerHealthInfo::new(0);
assert_eq!(health.crash_count, 0);
health.record_crash();
assert_eq!(health.crash_count, 1);
health.record_crash();
assert_eq!(health.crash_count, 2);
}
#[test]
fn test_worker_task_shutdown_variants() {
let graceful = WorkerTask::Shutdown { immediate: false };
let immediate = WorkerTask::Shutdown { immediate: true };
match graceful {
WorkerTask::Shutdown { immediate: false } => {}
_ => panic!("graceful shutdown should have immediate=false"),
}
match immediate {
WorkerTask::Shutdown { immediate: true } => {}
_ => panic!("immediate shutdown should have immediate=true"),
}
}
#[test]
fn test_worker_state_variants() {
assert_eq!(WorkerState::Idle, WorkerState::Idle);
assert_eq!(WorkerState::Processing, WorkerState::Processing);
assert_eq!(WorkerState::Stopping, WorkerState::Stopping);
assert_eq!(WorkerState::Stopped, WorkerState::Stopped);
assert_ne!(WorkerState::Idle, WorkerState::Processing);
assert_ne!(WorkerState::Stopping, WorkerState::Stopped);
}
#[tokio::test(flavor = "multi_thread")]
async fn test_worker_loop_exits_immediately_when_not_running() {
let queue = Arc::new(PriorityRequestQueue::new(100));
let response_channel = Arc::new(ResponseChannel::new());
let config = WorkerConfig::default();
let engine: Arc<RwLock<dyn InferenceEngine + Send + Sync>> =
Arc::new(RwLock::new(MockEngine));
let service = Arc::new(RwLock::new(EmbeddingService::new(engine, None)));
let manager = WorkerManager::new(queue, response_channel, config, service);
manager.running.store(false, Ordering::SeqCst);
manager.spawn_worker().await;
let deadline = tokio::time::Instant::now() + Duration::from_secs(5);
loop {
if manager.current_workers() == 0 {
break;
}
if tokio::time::Instant::now() >= deadline {
panic!(
"worker did not exit within 5s when running=false (current_workers={})",
manager.current_workers()
);
}
tokio::time::sleep(Duration::from_millis(50)).await;
}
let health = manager.worker_health.lock().await;
assert_eq!(health.len(), 1);
assert!(!health[0].is_alive, "worker must be marked dead after exit");
}
#[tokio::test(flavor = "multi_thread")]
async fn test_worker_loop_ignores_process_request_via_task_receiver() {
let queue = Arc::new(PriorityRequestQueue::new(100));
let response_channel = Arc::new(ResponseChannel::new());
let config = WorkerConfig::default();
let engine: Arc<RwLock<dyn InferenceEngine + Send + Sync>> =
Arc::new(RwLock::new(MockEngine));
let service = Arc::new(RwLock::new(EmbeddingService::new(engine, None)));
let manager = WorkerManager::new(queue, response_channel, config, service);
manager.spawn_worker().await;
assert_eq!(manager.current_workers(), 1);
{
let senders = manager.worker_senders.lock().await;
senders[0]
.send(WorkerTask::ProcessRequest {
request_id: "ignored-1".to_string(),
embed_request: EmbedRequest {
text: "hello".to_string(),
normalize: Some(true),
},
})
.await
.expect("send ProcessRequest must succeed");
}
tokio::time::sleep(Duration::from_millis(200)).await;
assert_eq!(
manager.current_workers(),
1,
"worker must still be running after receiving ProcessRequest"
);
{
let senders = manager.worker_senders.lock().await;
senders[0]
.send(WorkerTask::Shutdown { immediate: true })
.await
.unwrap();
}
let deadline = tokio::time::Instant::now() + Duration::from_secs(2);
loop {
if manager.current_workers() == 0 {
break;
}
if tokio::time::Instant::now() >= deadline {
panic!("worker did not exit after Shutdown");
}
tokio::time::sleep(Duration::from_millis(50)).await;
}
}
#[tokio::test]
async fn test_spawn_worker_appends_sender_and_health() {
let queue = Arc::new(PriorityRequestQueue::new(100));
let response_channel = Arc::new(ResponseChannel::new());
let config = WorkerConfig::default();
let engine: Arc<RwLock<dyn InferenceEngine + Send + Sync>> =
Arc::new(RwLock::new(MockEngine));
let service = Arc::new(RwLock::new(EmbeddingService::new(engine, None)));
let manager = WorkerManager::new(queue, response_channel, config, service);
manager.spawn_worker().await;
manager.spawn_worker().await;
manager.spawn_worker().await;
assert_eq!(manager.current_workers(), 3);
assert_eq!(manager.worker_senders.lock().await.len(), 3);
assert_eq!(manager.worker_health.lock().await.len(), 3);
let health = manager.worker_health.lock().await;
let mut ids: Vec<usize> = health.iter().map(|h| h.worker_id).collect();
ids.sort();
assert_eq!(ids, vec![0, 1, 2]);
manager.running.store(false, Ordering::SeqCst);
let senders = manager.worker_senders.lock().await.clone();
for s in &senders {
let _ = s.send(WorkerTask::Shutdown { immediate: true }).await;
}
tokio::time::sleep(Duration::from_millis(100)).await;
}
#[tokio::test(flavor = "multi_thread")]
async fn test_worker_loop_processes_multiple_requests() {
let queue = Arc::new(PriorityRequestQueue::new(100));
let response_channel = Arc::new(ResponseChannel::new());
let config = WorkerConfig::default();
let engine: Arc<RwLock<dyn InferenceEngine + Send + Sync>> =
Arc::new(RwLock::new(MockEngine));
let service = Arc::new(RwLock::new(EmbeddingService::new(engine, None)));
let manager = WorkerManager::new(
Arc::clone(&queue),
Arc::clone(&response_channel),
config,
service,
);
let mut rxs = Vec::new();
for i in 0..5 {
let req_id = format!("multi-{}", i);
rxs.push((i, response_channel.register(req_id).await));
}
for i in 0..5 {
let (tx, _) = tokio::sync::oneshot::channel();
let request = QueuedRequest {
request_id: format!("multi-{}", i),
embed_request: EmbedRequest {
text: format!("text-{}", i),
normalize: Some(true),
},
priority: Priority::Normal,
submitted_at: std::time::Instant::now(),
timeout: Duration::from_secs(30),
source: RequestSource::http("127.0.0.1".to_string()),
response_tx: tx,
};
queue.enqueue(request).await.unwrap();
}
manager.spawn_worker().await;
for (i, rx) in rxs {
let result = tokio::time::timeout(Duration::from_secs(5), rx).await;
assert!(result.is_ok(), "response {} should arrive within 5s", i);
let response_result = result.unwrap().unwrap();
assert!(response_result.is_ok());
assert_eq!(response_result.unwrap().dimension, 8);
}
manager.running.store(false, Ordering::SeqCst);
let senders = manager.worker_senders.lock().await.clone();
for s in &senders {
let _ = s.send(WorkerTask::Shutdown { immediate: true }).await;
}
tokio::time::sleep(Duration::from_millis(100)).await;
}
#[tokio::test(flavor = "current_thread")]
async fn test_worker_loop_updates_activity_time() {
let queue = Arc::new(PriorityRequestQueue::new(100));
let response_channel = Arc::new(ResponseChannel::new());
let config = WorkerConfig::default();
let engine: Arc<RwLock<dyn InferenceEngine + Send + Sync>> =
Arc::new(RwLock::new(MockEngine));
let service = Arc::new(RwLock::new(EmbeddingService::new(engine, None)));
let manager = WorkerManager::new(queue, response_channel, config, service);
manager.spawn_worker().await;
let initial_time = {
let health = manager.worker_health.lock().await;
health[0].last_active_time
};
let deadline = tokio::time::Instant::now() + Duration::from_secs(30);
#[allow(unused_assignments)]
let mut updated_time = initial_time;
loop {
{
let health = manager.worker_health.lock().await;
updated_time = health[0].last_active_time;
}
if updated_time > initial_time {
break;
}
if tokio::time::Instant::now() >= deadline {
panic!(
"last_active_time was not refreshed within 30s \
(initial={:?}, current={:?})",
initial_time, updated_time
);
}
tokio::time::sleep(Duration::from_millis(50)).await;
}
assert!(
updated_time > initial_time,
"last_active_time must be refreshed during worker loop"
);
manager.running.store(false, Ordering::SeqCst);
let senders = manager.worker_senders.lock().await.clone();
for s in &senders {
let _ = s.send(WorkerTask::Shutdown { immediate: true }).await;
}
let stop_deadline = tokio::time::Instant::now() + Duration::from_secs(5);
loop {
let all_stopped = {
let health = manager.worker_health.lock().await;
health.iter().all(|h| !h.is_alive)
};
if all_stopped {
break;
}
if tokio::time::Instant::now() >= stop_deadline {
break;
}
tokio::time::sleep(Duration::from_millis(50)).await;
}
}
#[tokio::test]
async fn test_process_request_with_none_normalize() {
let engine: Arc<RwLock<dyn InferenceEngine + Send + Sync>> =
Arc::new(RwLock::new(MockEngine));
let service = Arc::new(RwLock::new(EmbeddingService::new(engine, None)));
let (tx, _rx) = tokio::sync::oneshot::channel();
let request = QueuedRequest {
request_id: "test-none-norm".to_string(),
embed_request: EmbedRequest {
text: "hello".to_string(),
normalize: None,
},
priority: Priority::Normal,
submitted_at: std::time::Instant::now(),
timeout: Duration::from_secs(30),
source: RequestSource::http("127.0.0.1".to_string()),
response_tx: tx,
};
let result = WorkerManager::process_request(&request, &service).await;
assert!(result.is_ok());
assert_eq!(result.unwrap().dimension, 8);
}
#[tokio::test(flavor = "multi_thread")]
async fn test_shutdown_with_no_workers_is_safe() {
let queue = Arc::new(PriorityRequestQueue::new(100));
let response_channel = Arc::new(ResponseChannel::new());
let config = WorkerConfig::default();
let engine: Arc<RwLock<dyn InferenceEngine + Send + Sync>> =
Arc::new(RwLock::new(MockEngine));
let service = Arc::new(RwLock::new(EmbeddingService::new(engine, None)));
let manager = WorkerManager::new(queue, response_channel, config, service);
assert_eq!(manager.current_workers(), 0);
tokio::time::timeout(Duration::from_secs(15), manager.shutdown())
.await
.expect("shutdown with no workers should complete");
assert!(!manager.running.load(Ordering::SeqCst));
}
#[tokio::test]
async fn test_spawn_single_worker_records_health_with_correct_id() {
let queue = Arc::new(PriorityRequestQueue::new(100));
let response_channel = Arc::new(ResponseChannel::new());
let config = WorkerConfig::default();
let engine: Arc<RwLock<dyn InferenceEngine + Send + Sync>> =
Arc::new(RwLock::new(MockEngine));
let service = Arc::new(RwLock::new(EmbeddingService::new(engine, None)));
let manager = WorkerManager::new(queue, response_channel, config, service);
manager.running.store(false, Ordering::SeqCst);
WorkerManager::spawn_single_worker(
&manager.current_workers,
&manager.worker_senders,
&manager.worker_health,
&manager.queue,
&manager.response_channel,
&manager.embedding_service,
&manager.config,
&manager.running,
)
.await;
let deadline = tokio::time::Instant::now() + Duration::from_secs(2);
loop {
if manager.current_workers() == 0 {
break;
}
if tokio::time::Instant::now() >= deadline {
panic!("worker did not exit within 2s");
}
tokio::time::sleep(Duration::from_millis(20)).await;
}
let health = manager.worker_health.lock().await;
assert_eq!(health.len(), 1);
assert_eq!(health[0].worker_id, 0);
assert!(!health[0].is_alive);
}
#[tokio::test(flavor = "multi_thread")]
async fn test_scaling_monitor_scales_up_workers() {
let queue = Arc::new(PriorityRequestQueue::new(500));
let response_channel = Arc::new(ResponseChannel::new());
let config = WorkerConfig {
min_workers: 0,
max_workers: 4,
scale_up_threshold: 10,
scale_down_threshold: 5,
idle_timeout_secs: 60,
scale_check_interval_secs: 1,
};
let engine: Arc<RwLock<dyn InferenceEngine + Send + Sync>> =
Arc::new(RwLock::new(MockEngine));
let service = Arc::new(RwLock::new(EmbeddingService::new(engine, None)));
let manager = WorkerManager::new(
Arc::clone(&queue),
Arc::clone(&response_channel),
config,
service,
);
for i in 0..200 {
let (tx, _) = tokio::sync::oneshot::channel();
let request = QueuedRequest {
request_id: format!("scale-up-{}", i),
embed_request: EmbedRequest {
text: format!("text-{}", i),
normalize: Some(true),
},
priority: Priority::Normal,
submitted_at: std::time::Instant::now(),
timeout: Duration::from_secs(30),
source: RequestSource::http("127.0.0.1".to_string()),
response_tx: tx,
};
queue.enqueue(request).await.unwrap();
}
manager.start().await.unwrap();
assert_eq!(manager.current_workers(), 0, "start() spawns min_workers=0");
let deadline = tokio::time::Instant::now() + Duration::from_secs(10);
loop {
if manager.current_workers() > 0 {
break;
}
if tokio::time::Instant::now() >= deadline {
panic!(
"scaling monitor did not scale up within 10s (current_workers={})",
manager.current_workers()
);
}
tokio::time::sleep(Duration::from_millis(100)).await;
}
assert!(
manager.current_workers() > 0,
"workers should have been scaled up"
);
manager.running.store(false, Ordering::SeqCst);
let senders = manager.worker_senders.lock().await.clone();
for s in &senders {
let _ = s.send(WorkerTask::Shutdown { immediate: true }).await;
}
let stop_deadline = tokio::time::Instant::now() + Duration::from_secs(5);
loop {
if manager.current_workers() == 0 {
break;
}
if tokio::time::Instant::now() >= stop_deadline {
break;
}
tokio::time::sleep(Duration::from_millis(50)).await;
}
}
#[tokio::test(flavor = "multi_thread")]
async fn test_scaling_monitor_scales_down_workers() {
let queue = Arc::new(PriorityRequestQueue::new(100));
let response_channel = Arc::new(ResponseChannel::new());
let config = WorkerConfig {
min_workers: 1,
max_workers: 4,
scale_up_threshold: 100,
scale_down_threshold: 10,
idle_timeout_secs: 60,
scale_check_interval_secs: 1,
};
let engine: Arc<RwLock<dyn InferenceEngine + Send + Sync>> =
Arc::new(RwLock::new(MockEngine));
let service = Arc::new(RwLock::new(EmbeddingService::new(engine, None)));
let manager = WorkerManager::new(
Arc::clone(&queue),
Arc::clone(&response_channel),
config,
service,
);
manager.start().await.unwrap();
assert_eq!(manager.current_workers(), 1, "start() spawns min_workers=1");
manager.spawn_worker().await;
manager.spawn_worker().await;
assert_eq!(manager.current_workers(), 3);
let deadline = tokio::time::Instant::now() + Duration::from_secs(10);
loop {
if manager.current_workers() < 3 {
break;
}
if tokio::time::Instant::now() >= deadline {
panic!(
"scaling monitor did not scale down within 10s (current_workers={})",
manager.current_workers()
);
}
tokio::time::sleep(Duration::from_millis(100)).await;
}
assert!(
manager.current_workers() < 3,
"workers should have been scaled down"
);
manager.running.store(false, Ordering::SeqCst);
let senders = manager.worker_senders.lock().await.clone();
for s in &senders {
let _ = s.send(WorkerTask::Shutdown { immediate: true }).await;
}
let stop_deadline = tokio::time::Instant::now() + Duration::from_secs(5);
loop {
if manager.current_workers() == 0 {
break;
}
if tokio::time::Instant::now() >= stop_deadline {
break;
}
tokio::time::sleep(Duration::from_millis(50)).await;
}
}
#[tokio::test(flavor = "multi_thread")]
async fn test_worker_loop_enters_idle_backoff_when_channel_closed() {
let queue = Arc::new(PriorityRequestQueue::new(100));
let response_channel = Arc::new(ResponseChannel::new());
let config = WorkerConfig {
min_workers: 0,
max_workers: 2,
..Default::default()
};
let engine: Arc<RwLock<dyn InferenceEngine + Send + Sync>> =
Arc::new(RwLock::new(MockEngine));
let service = Arc::new(RwLock::new(EmbeddingService::new(engine, None)));
let manager = WorkerManager::new(queue, response_channel, config, service);
manager.spawn_worker().await;
assert_eq!(manager.current_workers(), 1);
{
let mut senders = manager.worker_senders.lock().await;
senders.clear();
}
tokio::time::sleep(Duration::from_millis(300)).await;
assert_eq!(
manager.current_workers(),
1,
"worker should still be running after idle backoff"
);
manager.running.store(false, Ordering::SeqCst);
let deadline = tokio::time::Instant::now() + Duration::from_secs(10);
loop {
if manager.current_workers() == 0 {
break;
}
if tokio::time::Instant::now() >= deadline {
panic!(
"worker did not exit within 10s after running=false (current_workers={})",
manager.current_workers()
);
}
tokio::time::sleep(Duration::from_millis(100)).await;
}
}
}