use crate::EmbeddingEngine;
use crate::server::channel::{TextInput, WorkerRequest, WorkerResponse};
use std::sync::{Arc, Mutex};
use std::thread;
use std::time::{Duration, Instant};
use tokio::sync::mpsc;
use tracing::{debug, error, info, warn};
#[derive(Debug, Clone)]
pub struct InferenceConfig {
pub max_batch_size: usize,
pub batch_timeout_ms: u64,
}
impl Default for InferenceConfig {
fn default() -> Self {
Self {
max_batch_size: 8,
batch_timeout_ms: 10,
}
}
}
pub struct InferenceWorker {
engine: Arc<Mutex<EmbeddingEngine>>,
receiver: mpsc::Receiver<WorkerRequest>,
config: InferenceConfig,
}
impl InferenceWorker {
pub fn new(
engine: Arc<Mutex<EmbeddingEngine>>,
receiver: mpsc::Receiver<WorkerRequest>,
config: InferenceConfig,
) -> Self {
Self {
engine,
receiver,
config,
}
}
pub fn run(mut self) {
info!("Inference worker starting");
loop {
let batch = self.collect_batch();
if batch.is_empty() {
break;
}
debug!("Processing batch of {} requests", batch.len());
let start = Instant::now();
self.process_batch(batch);
let elapsed = start.elapsed();
debug!("Batch processed in {:?}", elapsed);
}
info!("Inference worker shutting down");
}
fn collect_batch(&mut self) -> Vec<WorkerRequest> {
let mut batch = Vec::new();
let batch_start = Instant::now();
let timeout = Duration::from_millis(self.config.batch_timeout_ms);
match self.receiver.blocking_recv() {
Some(request) => batch.push(request),
None => return batch, }
while batch.len() < self.config.max_batch_size {
let elapsed = batch_start.elapsed();
if elapsed >= timeout {
break;
}
if let Ok(request) = self.receiver.try_recv() {
batch.push(request);
} else {
std::thread::sleep(Duration::from_micros(100));
if batch_start.elapsed() >= timeout {
break;
}
}
}
batch
}
fn process_batch(&self, batch: Vec<WorkerRequest>) {
let model_name = &batch[0].model;
let mut all_texts = Vec::new();
let mut request_metadata = Vec::new();
for request in &batch {
let start_idx = all_texts.len();
match &request.input {
TextInput::Single(text) => {
all_texts.push(text.as_str());
request_metadata.push((start_idx, 1));
}
TextInput::Batch(texts) => {
let count = texts.len();
for text in texts {
all_texts.push(text.as_str());
}
request_metadata.push((start_idx, count));
}
}
}
debug!(
"Processing {} texts across {} requests",
all_texts.len(),
batch.len()
);
let result = {
let Ok(engine) = self.engine.lock() else {
error!("Engine lock poisoned during batch processing");
for request in batch {
drop(request.response_tx);
}
return;
};
engine.embed_batch(Some(model_name), &all_texts)
};
match result {
Ok(all_embeddings) => {
for (request, (start_idx, count)) in batch.into_iter().zip(request_metadata.iter())
{
let request_embeddings =
all_embeddings[*start_idx..*start_idx + *count].to_vec();
let response = WorkerResponse {
embeddings: request_embeddings.clone(),
token_count: request_embeddings.iter().map(Vec::len).sum::<usize>() / 10, processing_time_ms: 0, };
if request.response_tx.send(response).is_err() {
warn!(
"Failed to send response for request {} (client may have timed out)",
request.id
);
}
}
}
Err(e) => {
error!("Batch processing failed: {}", e);
for request in batch {
let response = WorkerResponse {
embeddings: vec![],
token_count: 0,
processing_time_ms: 0,
};
let _ = request.response_tx.send(response);
}
}
}
}
pub fn spawn(
engine: Arc<Mutex<EmbeddingEngine>>,
receiver: mpsc::Receiver<WorkerRequest>,
config: InferenceConfig,
) -> thread::JoinHandle<()> {
thread::spawn(move || {
let worker = InferenceWorker::new(engine, receiver, config);
worker.run();
})
}
}