Backend de alta concurrencia para orquestar inferencias de LLM: cola con backpressure, pool de workers con límite de GPU slots, métricas broadcast en tiempo real y shutdown graceful.
// Buffer de 50: bloquea productores si está lleno let (request_tx, request_rx) = mpsc::channel(50); // Envío no bloqueante → error si llena request_tx .try_send((req, reply_tx)) .map_err(|_| ServiceError::QueueFull(50))?; // Recepción en el dispatcher while let Some((req, tx)) = rx.recv().await { // despachar al worker... }
let mut set = JoinSet::new(); // Lanzar 8 inferencias concurrentes for req in requests { let e = engine.clone(); set.spawn(async move { e.infer(req).await }); } // Recoger resultados según completan while let Some(res) = set.join_next().await { match res.unwrap() { Ok(resp) => info!("OK: {}", resp.text), Err(e) => error!("Fallo: {}", e), } }
// El dispatcher reacciona a cualquier evento loop { select! { _ = cancel.cancelled() => { info!("Shutdown → dreno tareas"); break; } Some((req, tx)) = rx.recv() => { // nueva petición → spawn worker set.spawn(run_inference(req)); } Some(r) = set.join_next() => { // tarea terminada → libera handle } } }
let semaphore = Arc::new(Semaphore::new(4)); // GPU slots // Adquirir permiso antes de inferir let permit = semaphore .clone() .acquire_owned() .await .unwrap(); tokio::spawn!(async move { run_inference(req).await; drop(permit); // libera slot de GPU });
let token = CancellationToken::new(); let child = token.clone(); // Worker respeta la cancelación select! { _ = child.cancelled() => { info!("Worker detenido limpiamente"); return; } r = do_inference() => { /* procesar */ } } // Desde el motor → cancela todos pub async fn shutdown(&self) { token.cancel(); // instantáneo y thread-safe }
let (tx, _) = broadcast::channel(256); // Emitir desde cada inferencia tx.send(MetricEvent { latency_ms, success }); // Agregar con streams por lotes de 10 let mut chunks = stream.chunks(10); while let Some(batch) = chunks.next().await { let p50 = latencies[total / 2]; let p95 = latencies[total * 95 / 100]; info!(p50_ms = p50, p95_ms = p95, "Batch"); }
#[async_trait] pub trait ModelBackend: Send + Sync { fn model_id(&self) -> &str; async fn generate( &self, prompt: &str, max_tokens: u32 ) -> Result<(String, u32), ServiceError>; } // Implementar para cada modelo #[async_trait] impl ModelBackend for MockLlmBackend { ... }
#[derive(Error, Debug)] pub enum ServiceError { #[error("Timeout after {0:?}")] Timeout(Duration), #[error("Model '{0}' not found")] ModelNotFound(String), #[error("Queue full ({0})")] QueueFull(usize), #[error("Shutting down")] ShuttingDown, }
[dependencies] tokio = { version = "1", features = ["full"] } futures = "0.3" async-trait = "0.1" thiserror = "1.0" anyhow = "1.0" tracing = "0.1" tracing-subscriber = "0.3" tokio-util = { version = "0.7", features = ["sync"] }
#[tokio::test] async fn test_concurrent_inferences() { let mut set = JoinSet::new(); for i in 0..8 { set.spawn(engine.infer(req(i))); } let mut ok = 0; while let Some(r) = set.join_next().await { if r.unwrap().is_ok() { ok += 1; } } assert_eq!(ok, 8); }