use crate::EmbeddingEngine;
use crate::server::channel::{TextInput, WorkerRequest, WorkerResponse};
use std::sync::{Arc, Mutex};
use std::thread;
use std::time::Instant;
use tokio::sync::mpsc;
use tracing::{debug, error, info, warn};
pub struct Worker {
id: usize,
engine: Arc<Mutex<EmbeddingEngine>>,
receiver: mpsc::Receiver<WorkerRequest>,
}
impl Worker {
pub fn new(
id: usize,
engine: Arc<Mutex<EmbeddingEngine>>,
receiver: mpsc::Receiver<WorkerRequest>,
) -> Self {
Self {
id,
engine,
receiver,
}
}
pub fn run(mut self) {
info!("Worker {} starting", self.id);
while let Some(request) = self.receiver.blocking_recv() {
debug!("Worker {} processing request {:?}", self.id, request.id);
let start = Instant::now();
let result = {
let Ok(engine) = self.engine.lock() else {
error!("Worker {} engine lock poisoned", self.id);
drop(request.response_tx);
continue;
};
match &request.input {
TextInput::Single(text) => {
engine
.embed(Some(&request.model), text)
.map(|embedding| vec![embedding])
}
TextInput::Batch(texts) => {
let text_refs: Vec<&str> =
texts.iter().map(std::string::String::as_str).collect();
engine.embed_batch(Some(&request.model), &text_refs)
}
}
};
let response = match result {
Ok(embeddings) => {
let token_count = embeddings.iter().map(std::vec::Vec::len).sum::<usize>() / 10;
WorkerResponse {
embeddings,
token_count,
processing_time_ms: u64::try_from(start.elapsed().as_millis())
.unwrap_or(u64::MAX),
}
}
Err(e) => {
error!("Worker {} failed to generate embeddings: {}", self.id, e);
WorkerResponse {
embeddings: vec![],
token_count: 0,
processing_time_ms: u64::try_from(start.elapsed().as_millis())
.unwrap_or(u64::MAX),
}
}
};
if request.response_tx.send(response).is_err() {
warn!(
"Worker {} failed to send response for request {:?} (client may have timed out)",
self.id, request.id
);
}
}
info!("Worker {} shutting down", self.id);
}
pub fn spawn(
id: usize,
engine: Arc<Mutex<EmbeddingEngine>>,
receiver: mpsc::Receiver<WorkerRequest>,
) -> thread::JoinHandle<()> {
thread::spawn(move || {
let worker = Worker::new(id, engine, receiver);
worker.run();
})
}
}