use crate::EmbeddingEngine;
use crate::server::channel::WorkerRequest;
use crate::server::inference_worker::{InferenceConfig, InferenceWorker};
use std::sync::Arc;
use std::thread;
use tokio::sync::mpsc;
use tracing::{debug, info, warn};
pub struct Dispatcher {
inference_tx: mpsc::Sender<WorkerRequest>,
#[allow(dead_code)]
handle: Option<thread::JoinHandle<()>>,
}
impl Clone for Dispatcher {
fn clone(&self) -> Self {
Self {
inference_tx: self.inference_tx.clone(),
handle: None, }
}
}
impl Dispatcher {
pub fn new(n_seq_max: usize, queue_size: usize) -> Self {
info!(
"Creating dispatcher with inference worker (n_seq_max={})",
n_seq_max
);
let engine = EmbeddingEngine::instance()
.expect("EmbeddingEngine should be initialized before creating Dispatcher");
let (tx, rx) = mpsc::channel::<WorkerRequest>(queue_size);
let config = InferenceConfig {
max_batch_size: n_seq_max,
batch_timeout_ms: 10, };
let handle = InferenceWorker::spawn(Arc::clone(&engine), rx, config);
info!("Dispatcher created with dedicated inference worker");
Self {
inference_tx: tx,
handle: Some(handle),
}
}
pub async fn send(&self, request: WorkerRequest) -> Result<(), String> {
debug!("Routing request {:?} to inference worker", request.id);
self.inference_tx.send(request).await.map_err(|e| {
warn!("Inference worker channel full or closed: {}", e);
format!("Failed to send request to inference worker: {e}")
})
}
pub fn is_ready(&self) -> bool {
!self.inference_tx.is_closed()
}
pub fn worker_count(&self) -> usize {
1
}
pub fn shutdown(self) {
info!("Shutting down dispatcher and inference worker");
drop(self);
info!("Inference worker signaled to shutdown");
}
}