#![allow(
dead_code,
reason = "WorkerManager is used via queue in handler; tests cover all methods, production uses shared queue"
)]
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, QueuedRequest, ServiceRequest};
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,
request: ServiceRequest,
},
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 async fn assemble_batch<F, Fut>(
first: QueuedRequest,
mut try_dequeue: F,
max_batch_size: usize,
batch_wait_ms: u64,
) -> Vec<QueuedRequest>
where
F: FnMut() -> Fut,
Fut: std::future::Future<Output = Option<QueuedRequest>>,
{
let mut batch = vec![first];
let cap = max_batch_size.max(1);
if batch.len() >= cap {
return batch;
}
if batch_wait_ms == 0 {
return batch;
}
let deadline = tokio::time::Instant::now() + Duration::from_millis(batch_wait_ms);
loop {
if batch.len() >= cap {
break;
}
let now = tokio::time::Instant::now();
if now >= deadline {
break;
}
match try_dequeue().await {
Some(req) => {
batch.push(req);
}
None => {
let remaining = deadline.saturating_duration_since(now);
let sleep_for = std::cmp::min(remaining, Duration::from_millis(1));
if sleep_for.is_zero() {
break;
}
tokio::time::sleep(sleep_for).await;
}
}
}
batch
}
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>>>,
bg_tasks: Arc<Mutex<tokio::task::JoinSet<()>>>,
}
#[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 {
let max_workers = config.max_workers;
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::with_capacity(max_workers))),
embedding_service,
worker_health: Arc::new(Mutex::new(Vec::with_capacity(max_workers))),
bg_tasks: Arc::new(Mutex::new(tokio::task::JoinSet::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;
}
{
let mut tasks = self.bg_tasks.lock().await;
tasks.abort_all();
}
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,
&self.bg_tasks,
)
.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>,
bg_tasks: &Arc<Mutex<tokio::task::JoinSet<()>>>,
) {
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);
let handle = tokio::spawn(async move {
Self::worker_loop(
worker_id,
task_receiver,
queue,
response_channel,
config,
running,
embedding_service,
worker_health,
current_workers,
)
.await;
});
bg_tasks.lock().await.spawn(async move {
let _ = handle.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;
let wait_start = tokio::time::Instant::now();
let batch = assemble_batch(
request,
|| queue.dequeue(),
config.max_batch_size,
config.batch_wait_ms,
)
.await;
let waited_secs = wait_start.elapsed().as_secs_f64();
#[cfg(feature = "http")]
if let Some(collector) = crate::metrics::prometheus_exporter::global_collector()
{
collector.observe_batch("embed", batch.len(), waited_secs);
}
debug!(
"Worker {} assembled batch of {} (wait {:.3}s, window {}ms)",
worker_id,
batch.len(),
waited_secs,
config.batch_wait_ms
);
const QUEUE_EXPIRY: Duration = Duration::from_secs(30);
let now = std::time::Instant::now();
let mut valid_batch = Vec::with_capacity(batch.len());
for req in batch {
if now.duration_since(req.submitted_at) > QUEUE_EXPIRY {
warn!(
"Request {} expired in queue ({:.1}s), rejecting",
req.request_id,
now.duration_since(req.submitted_at).as_secs_f64()
);
response_channel
.complete(
req.request_id.clone(),
Err(VecboostError::RateLimitExceeded(
"Request expired in queue".to_string(),
)),
)
.await;
} else {
valid_batch.push(req);
}
}
if valid_batch.is_empty() {
continue;
}
debug!(
"Worker {} processing batch of {} requests",
worker_id, valid_batch.len()
);
Self::process_batch_requests(&valid_batch, &embedding_service, &response_channel).await;
}
result = tokio::time::timeout(
Duration::from_secs(config.idle_timeout_secs),
queue.notify().notified(),
) => {
match result {
Ok(()) => {
debug!("Worker {} notified of new request", worker_id);
idle_count = 0;
}
Err(_) => {
idle_count = idle_count.saturating_add(1);
debug!(
"Worker {} idle timeout ({}s), idle_count={}",
worker_id, config.idle_timeout_secs, idle_count
);
}
}
}
}
if idle_count > MAX_IDLE_COUNT && queue.size() == 0 && 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 = match &request.request {
ServiceRequest::Embed(req) => req,
ServiceRequest::Rerank(_) => {
return Err(VecboostError::InternalError(
"Rerank not supported by embedding worker".to_string(),
));
}
};
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 process_batch_requests(
batch: &[super::queue::QueuedRequest],
embedding_service: &Arc<RwLock<EmbeddingService>>,
response_channel: &Arc<ResponseChannel>,
) {
if batch.is_empty() {
return;
}
if batch.len() == 1 {
let result = Self::process_request(&batch[0], embedding_service).await;
response_channel
.complete(batch[0].request_id.clone(), result)
.await;
return;
}
let mut texts = Vec::with_capacity(batch.len());
let mut normalize_flags = Vec::with_capacity(batch.len());
let mut valid_indices = Vec::with_capacity(batch.len());
for (i, req) in batch.iter().enumerate() {
match &req.request {
ServiceRequest::Embed(embed_req) => {
texts.push(embed_req.text.clone());
normalize_flags.push(embed_req.normalize.unwrap_or(false));
valid_indices.push(i);
}
ServiceRequest::Rerank(_) => {
response_channel
.complete(
req.request_id.clone(),
Err(VecboostError::InternalError(
"Rerank not supported by embedding worker".to_string(),
)),
)
.await;
}
}
}
if texts.is_empty() {
return;
}
let service_guard = embedding_service.read().await;
let batch_started = std::time::Instant::now();
let batch_result = service_guard.embed_batch_texts(&texts).await;
drop(service_guard);
match batch_result {
Ok(embeddings) => {
let batch_millis = batch_started.elapsed().as_millis();
for (j, &idx) in valid_indices.iter().enumerate() {
let req = &batch[idx];
if j < embeddings.len() {
let mut embedding = embeddings[j].clone();
if normalize_flags[idx] {
crate::utils::vector::normalize_l2(&mut embedding).ok();
}
let dimension = embedding.len();
response_channel
.complete(
req.request_id.clone(),
Ok(EmbedResponse {
dimension,
embedding,
processing_time_ms: batch_millis,
information_retention_rate: None,
}),
)
.await;
} else {
response_channel
.complete(
req.request_id.clone(),
Err(VecboostError::InternalError(
"Batch inference returned fewer embeddings than inputs"
.to_string(),
)),
)
.await;
}
}
}
Err(e) => {
warn!("Batch inference failed: {}", e);
for req in batch.iter() {
if matches!(req.request, ServiceRequest::Embed(_)) {
response_channel
.complete(req.request_id.clone(), Err(e.clone()))
.await;
}
}
}
}
}
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);
let bg_tasks = Arc::clone(&self.bg_tasks);
let bg_tasks_for_spawn = Arc::clone(&bg_tasks);
bg_tasks_for_spawn.lock().await.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,
&bg_tasks,
)
.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;
use std::collections::VecDeque;
use std::sync::Arc as StdArc;
use tokio::sync::Mutex as TokioMutex;
fn make_queued(id: &str) -> QueuedRequest {
QueuedRequest {
request_id: id.to_string(),
request: ServiceRequest::Embed(EmbedRequest {
text: id.to_string(),
normalize: Some(false),
}),
priority: Priority::Normal,
submitted_at: std::time::Instant::now(),
timeout: Duration::from_secs(30),
source: RequestSource::http("127.0.0.1".to_string()),
}
}
#[tokio::test]
async fn test_assemble_batch_collects_arrivals_within_window() {
let pending: StdArc<TokioMutex<VecDeque<QueuedRequest>>> =
StdArc::new(TokioMutex::new(VecDeque::new()));
let pending_clone = StdArc::clone(&pending);
tokio::spawn(async move {
tokio::time::sleep(Duration::from_millis(5)).await;
pending_clone.lock().await.push_back(make_queued("req-2"));
tokio::time::sleep(Duration::from_millis(5)).await;
pending_clone.lock().await.push_back(make_queued("req-3"));
});
let first = make_queued("req-1");
let batch = assemble_batch(
first,
|| {
let pending = StdArc::clone(&pending);
async move { pending.lock().await.pop_front() }
},
8,
50,
)
.await;
assert_eq!(batch.len(), 3, "窗口内陆续到达 3 请求应单次组装返回 3 条");
assert_eq!(batch[0].request_id, "req-1");
}
#[tokio::test]
async fn test_assemble_batch_zero_wait_returns_first_only() {
let first = make_queued("only");
let batch = assemble_batch(first, || async { Some(make_queued("late")) }, 8, 0).await;
assert_eq!(batch.len(), 1, "batch_wait_ms=0 应立即返回仅首请求");
assert_eq!(batch[0].request_id, "only");
}
#[tokio::test]
async fn test_assemble_batch_full_returns_early() {
let first = make_queued("a");
let batch = assemble_batch(first, || async { Some(make_queued("extra")) }, 2, 100).await;
assert_eq!(batch.len(), 2, "凑满 max_batch_size 应提前返回");
}
#[tokio::test]
async fn test_process_batch_same_text_byte_equal_and_ordered() {
use crate::service::embedding::EmbeddingService;
let engine: Arc<RwLock<dyn InferenceEngine + Send + Sync>> =
Arc::new(RwLock::new(MockEngine));
let service = Arc::new(RwLock::new(EmbeddingService::new(engine, None)));
let channel = Arc::new(ResponseChannel::new());
let mk = |id: &str, text: &str| QueuedRequest {
request_id: id.to_string(),
request: ServiceRequest::Embed(EmbedRequest {
text: text.to_string(),
normalize: Some(false),
}),
priority: Priority::Normal,
submitted_at: std::time::Instant::now(),
timeout: Duration::from_secs(30),
source: RequestSource::http("127.0.0.1".to_string()),
};
let baseline = service
.read()
.await
.embed_batch_texts(&["hello world".to_string()])
.await
.expect("baseline embed must succeed");
let b1 = mk("b1", "hello world");
let b2 = mk("b2", "hello world");
let b3 = mk("b3", "other text");
let batch = vec![b1, b2, b3];
let rx1 = channel.register("b1".to_string()).await;
let rx2 = channel.register("b2".to_string()).await;
let rx3 = channel.register("b3".to_string()).await;
WorkerManager::process_batch_requests(&batch, &service, &channel).await;
let r1 = rx1.await.expect("b1 response").expect("b1 ok");
let r2 = rx2.await.expect("b2 response").expect("b2 ok");
let r3 = rx3.await.expect("b3 response").expect("b3 ok");
assert_eq!(
r1.embedding, baseline[0],
"批内结果须与同函数单文本基线字节等同"
);
assert_eq!(r1.embedding, r2.embedding, "相同文本批内 scatter 须一致");
assert_eq!(r1.dimension, r3.dimension);
}
struct MockEngine;
#[async_trait]
impl InferenceEngine for MockEngine {
fn embed(&self, _text: &str) -> Result<Vec<f32>, VecboostError> {
Ok(vec![1.0, 2.0, 3.0, 4.0, 5.0, 6.0, 7.0, 8.0])
}
fn embed_batch(&self, texts: &[String]) -> Result<Vec<Vec<f32>>, VecboostError> {
Ok(texts
.iter()
.map(|_| vec![1.0, 2.0, 3.0, 4.0, 5.0, 6.0, 7.0, 8.0])
.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(())
}
}
#[tokio::test]
async fn test_worker_manager_creation() {
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::test]
async fn test_worker_manager_start() {
let queue = Arc::new(PriorityRequestQueue::new(100));
let response_channel = Arc::new(ResponseChannel::new());
let config = WorkerConfig {
min_workers: 2,
max_workers: 4,
..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(), 2);
}
#[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);
let bg_tasks = Arc::new(Mutex::new(tokio::task::JoinSet::new()));
WorkerManager::spawn_single_worker(
¤t_workers,
&worker_senders,
&worker_health,
&queue_clone,
&response_channel_clone,
&embedding_service_clone,
&config_clone,
&running,
&bg_tasks,
)
.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",
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 request = QueuedRequest {
request_id: "test-process-1".to_string(),
request: ServiceRequest::Embed(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()),
};
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 request = QueuedRequest {
request_id: "test-process-err".to_string(),
request: ServiceRequest::Embed(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()),
};
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 request = QueuedRequest {
request_id: "test-loop-1".to_string(),
request: ServiceRequest::Embed(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()),
};
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 request = QueuedRequest {
request_id: "test-loop-err".to_string(),
request: ServiceRequest::Embed(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()),
};
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 {
let health = manager.worker_health.lock().await;
if health.len() == 1 && !health[0].is_alive && manager.current_workers() == 0 {
break;
}
drop(health);
if tokio::time::Instant::now() >= deadline {
let health = manager.worker_health.lock().await;
panic!(
"worker did not exit within 2s (current_workers={}, health_len={}, is_alive={})",
manager.current_workers(),
health.len(),
health.first().map(|h| h.is_alive).unwrap_or(false)
);
}
tokio::time::sleep(Duration::from_millis(50)).await;
}
}
#[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 {
let health = manager.worker_health.lock().await;
if health.len() == 1 && !health[0].is_alive && manager.current_workers() == 0 {
break;
}
drop(health);
if tokio::time::Instant::now() >= deadline {
let health = manager.worker_health.lock().await;
panic!(
"worker did not exit within 5s when running=false (current_workers={}, health_len={}, is_alive={})",
manager.current_workers(),
health.len(),
health.first().map(|h| h.is_alive).unwrap_or(false)
);
}
tokio::time::sleep(Duration::from_millis(50)).await;
}
}
#[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(),
request: ServiceRequest::Embed(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 request = QueuedRequest {
request_id: format!("multi-{}", i),
request: ServiceRequest::Embed(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()),
};
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 request = QueuedRequest {
request_id: "test-none-norm".to_string(),
request: ServiceRequest::Embed(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()),
};
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,
&manager.bg_tasks,
)
.await;
let deadline = tokio::time::Instant::now() + Duration::from_secs(2);
loop {
let health = manager.worker_health.lock().await;
if health.len() == 1
&& health[0].worker_id == 0
&& !health[0].is_alive
&& manager.current_workers() == 0
{
break;
}
drop(health);
if tokio::time::Instant::now() >= deadline {
let health = manager.worker_health.lock().await;
panic!(
"worker did not exit within 2s (current_workers={}, health_len={}, is_alive={})",
manager.current_workers(),
health.len(),
health.first().map(|h| h.is_alive).unwrap_or(false)
);
}
tokio::time::sleep(Duration::from_millis(20)).await;
}
}
#[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,
max_batch_size: 8,
batch_wait_ms: 5,
};
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 request = QueuedRequest {
request_id: format!("scale-up-{}", i),
request: ServiceRequest::Embed(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()),
};
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,
max_batch_size: 8,
batch_wait_ms: 5,
};
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_exits_after_idle_timeout_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,
idle_timeout_secs: 1, ..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(1500)).await;
assert_eq!(
manager.current_workers(),
1,
"worker should still be running after idle timeout"
);
manager.running.store(false, Ordering::SeqCst);
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 after running=false (current_workers={})",
manager.current_workers()
);
}
tokio::time::sleep(Duration::from_millis(100)).await;
}
}
}