use crate::engine::{ChatTurn, CompactionOutput, GenerationOutput, Generator}; use crate::metrics::{Metrics, WorkSource}; use crate::settings::{EngineSettings, TurnSettings}; use std::path::PathBuf; use std::sync::Arc; use std::sync::atomic::{AtomicBool, Ordering}; use std::sync::mpsc::{self, Receiver, Sender}; use std::thread; use std::time::{Duration, Instant}; const RUNTIME_PANIC_ERROR: &str = "The model runtime hit an internal error and was reset."; #[derive(Clone)] pub(crate) struct GenerationService { commands: Sender, metrics: Arc, } pub(crate) struct ActiveGeneration { pub(crate) events: Receiver, pub(crate) cancel: Arc, } impl Drop for ActiveGeneration { fn drop(&mut self) { self.cancel.store(true, Ordering::Relaxed); } } pub(crate) enum CheckpointTarget { Local { checkpoint: PathBuf, /// Shared cache for the deterministic rendered system prompt. bootstrap: Option, }, Transient(PathBuf), } pub(crate) enum GenerationEvent { Loading, Activity(&'static str), Chunk { reasoning: bool, content: String, }, Context { used: u32, limit: u32, tokens_per_second: Option, }, Finished(Result), Compacted(Result), Measured(Result), } enum Operation { Generate, Compact { reason: String, rebuild_system_prompt: String, }, Measure, } impl Operation { fn tracks_metrics(&self) -> bool { !matches!(self, Self::Measure) } fn error_handler(&self) -> fn(String) -> GenerationEvent { match self { Self::Generate => |error| GenerationEvent::Finished(Err(error)), Self::Compact { .. } => |error| GenerationEvent::Compacted(Err(error)), Self::Measure => |error| GenerationEvent::Measured(Err(error)), } } } struct Command { engine: EngineSettings, turn: TurnSettings, messages: Vec, checkpoint: Option, source: WorkSource, operation: Operation, idle_timeout: Duration, cancel: Arc, events: Sender, } pub(crate) struct CompactionInput { pub(crate) engine: EngineSettings, pub(crate) turn: TurnSettings, pub(crate) messages: Vec, pub(crate) reason: String, pub(crate) rebuild_system_prompt: String, pub(crate) checkpoint: PathBuf, pub(crate) idle_timeout: Duration, } struct RuntimeState { loaded: Option<(EngineSettings, Generator)>, last_used: Instant, idle_timeout: Duration, } impl GenerationService { pub(crate) fn spawn(metrics: Arc) -> Result { let (commands, receiver) = mpsc::channel::(); let worker_metrics = Arc::clone(&metrics); thread::Builder::new() .name("model-runtime".into()) .spawn(move || run(receiver, worker_metrics)) .map_err(|error| format!("Could not start the model runtime: {error}"))?; Ok(Self { commands, metrics }) } pub(crate) fn generate( &self, engine: EngineSettings, turn: TurnSettings, messages: Vec, checkpoint: CheckpointTarget, source: WorkSource, idle_timeout: Duration, ) -> Result { self.submit(CommandRequest { engine, turn, messages, checkpoint: Some(checkpoint), source, operation: Operation::Generate, idle_timeout, }) } pub(crate) fn compact(&self, input: CompactionInput) -> Result { self.submit(CommandRequest { engine: input.engine, turn: input.turn, messages: input.messages, checkpoint: Some(CheckpointTarget::Local { checkpoint: input.checkpoint, bootstrap: None, }), source: WorkSource::LocalChat, operation: Operation::Compact { reason: input.reason, rebuild_system_prompt: input.rebuild_system_prompt, }, idle_timeout: input.idle_timeout, }) } pub(crate) fn measure_context( &self, engine: EngineSettings, turn: TurnSettings, messages: Vec, idle_timeout: Duration, ) -> Result { self.submit(CommandRequest { engine, turn, messages, checkpoint: None, source: WorkSource::LocalChat, operation: Operation::Measure, idle_timeout, }) } fn submit(&self, request: CommandRequest) -> Result { let cancel = Arc::new(AtomicBool::new(false)); let (events, receiver) = mpsc::channel(); let tracked = request.operation.tracks_metrics(); if tracked { self.metrics.request_queued(request.source); } let command = Command { engine: request.engine, turn: request.turn, messages: request.messages, checkpoint: request.checkpoint, source: request.source, operation: request.operation, idle_timeout: request.idle_timeout, cancel: Arc::clone(&cancel), events, }; if self.commands.send(command).is_err() { if tracked { self.metrics.request_rejected(); } return Err("The model runtime stopped unexpectedly.".to_owned()); } Ok(ActiveGeneration { events: receiver, cancel, }) } } struct CommandRequest { engine: EngineSettings, turn: TurnSettings, messages: Vec, checkpoint: Option, source: WorkSource, operation: Operation, idle_timeout: Duration, } fn run(commands: Receiver, metrics: Arc) { let mut state = RuntimeState { loaded: None, last_used: Instant::now(), idle_timeout: Duration::from_secs(15 * 60), }; loop { match commands.recv_timeout(Duration::from_secs(1)) { Ok(command) => { let request_started = Instant::now(); let source = command.source; let events = command.events.clone(); let error_event = command.operation.error_handler(); let tracked = command.operation.tracks_metrics(); if tracked { metrics.request_started(source); } if let Err(error) = catch_runtime_panic(|| { run_command(command, &mut state, &metrics, request_started, source); }) { if tracked { metrics.request_failed(request_started.elapsed()); } if state.loaded.take().is_some() { metrics.unloaded(); } let _ = events.send(error_event(error)); } } Err(mpsc::RecvTimeoutError::Timeout) => { if state.loaded.is_some() && state.last_used.elapsed() >= state.idle_timeout { state.loaded = None; metrics.unloaded(); } } Err(mpsc::RecvTimeoutError::Disconnected) => break, } } } fn run_command( command: Command, state: &mut RuntimeState, metrics: &Arc, request_started: Instant, source: WorkSource, ) { let error_event = command.operation.error_handler(); let tracked = command.operation.tracks_metrics(); state.idle_timeout = command.idle_timeout; if command.cancel.load(Ordering::Relaxed) { if tracked { metrics.request_failed(request_started.elapsed()); } let _ = command .events .send(error_event("generation cancelled".into())); return; } if state .loaded .as_ref() .is_none_or(|(settings, _)| settings != &command.engine) { if state.loaded.take().is_some() { metrics.unloaded(); } let _ = command.events.send(GenerationEvent::Loading); metrics.loading(); let load_started = Instant::now(); state.loaded = match Generator::open(&command.engine, Arc::clone(metrics)) { Ok(generator) => { let summary = generator.summary(); metrics.loaded( summary.model, load_started.elapsed(), summary.mapped_bytes, summary.tensor_count, summary.vocabulary_size, ); Some((command.engine.clone(), generator)) } Err(error) => { if tracked { metrics.request_failed(request_started.elapsed()); } let _ = command.events.send(error_event(error)); None } }; } if let Some((_, generator)) = &mut state.loaded { let mut prefill_started = None::<(Instant, u32)>; let mut emit = |reasoning, content| { let _ = command .events .send(GenerationEvent::Chunk { reasoning, content }); }; let mut progress = |used, limit, tokens_per_second| { if let Some(speed) = tokens_per_second { metrics.generation_progress(used, limit, speed); } else { let (started, initial) = prefill_started.get_or_insert_with(|| (Instant::now(), used)); let elapsed = started.elapsed().as_secs_f32(); let speed = if elapsed > 0.0 { used.saturating_sub(*initial) as f32 / elapsed } else { 0.0 }; metrics.prefill_progress(used, limit, speed); } let _ = command.events.send(GenerationEvent::Context { used, limit, tokens_per_second, }); }; if let Operation::Compact { reason, rebuild_system_prompt, } = &command.operation { let Some(CheckpointTarget::Local { checkpoint, .. }) = &command.checkpoint else { unreachable!("compaction checkpoints are local") }; let result = generator.compact( &command.messages, &command.turn, rebuild_system_prompt, reason, checkpoint, &command.cancel, &mut progress, |activity| { let _ = command.events.send(GenerationEvent::Activity(activity)); }, ); match &result { Ok(_) => { metrics.request_finished(source, request_started.elapsed(), 0, 0, 0, None, 0) } Err(_) => metrics.request_failed(request_started.elapsed()), } let _ = command.events.send(GenerationEvent::Compacted(result)); state.last_used = Instant::now(); return; } if matches!(command.operation, Operation::Measure) { let result = generator.rendered_history_tokens(&command.messages, &command.turn); let _ = command.events.send(GenerationEvent::Measured(result)); state.last_used = Instant::now(); return; } let result = match command .checkpoint .expect("generation requires a checkpoint target") { CheckpointTarget::Local { checkpoint, bootstrap, } => generator.generate( &checkpoint, bootstrap.as_deref(), &command.messages, &command.turn, &command.cancel, &mut emit, &mut progress, |activity| { let _ = command.events.send(GenerationEvent::Activity(activity)); }, ), CheckpointTarget::Transient(directory) => generator.generate_transient( &directory, &command.messages, &command.turn, &command.cancel, &mut emit, &mut progress, ), }; match &result { Ok(output) => metrics.request_finished( source, request_started.elapsed(), output.prompt_tokens, output.cached_tokens, output.completion_tokens, output.previous_checkpoint_bytes, output.checkpoint_bytes, ), Err(_) => metrics.request_failed(request_started.elapsed()), } let _ = command.events.send(GenerationEvent::Finished(result)); state.last_used = Instant::now(); } } fn catch_runtime_panic(operation: impl FnOnce() -> T) -> Result { // Generator owns GPU handles and is not unwind-safe; callers discard it on error. std::panic::catch_unwind(std::panic::AssertUnwindSafe(operation)) .map_err(|_| RUNTIME_PANIC_ERROR.to_owned()) } #[cfg(test)] mod tests { use super::*; #[test] fn dropping_active_generation_cancels_it() { let cancel = Arc::new(AtomicBool::new(false)); let (_, events) = mpsc::channel(); drop(ActiveGeneration { events, cancel: Arc::clone(&cancel), }); assert!(cancel.load(Ordering::Relaxed)); } #[test] fn runtime_panic_recovery_accepts_the_next_operation() { assert_eq!( catch_runtime_panic(|| panic!("injected runtime panic")), Err(RUNTIME_PANIC_ERROR.to_owned()) ); assert_eq!(catch_runtime_panic(|| 42), Ok(42)); } #[test] fn operation_failures_use_the_matching_event() { assert!(matches!( Operation::Compact { reason: String::new(), rebuild_system_prompt: String::new(), } .error_handler()("stopped".into()), GenerationEvent::Compacted(Err(error)) if error == "stopped" )); assert!(matches!( Operation::Measure.error_handler()("stopped".into()), GenerationEvent::Measured(Err(error)) if error == "stopped" )); } #[test] fn operation_tracking_is_consistent() { assert!(Operation::Generate.tracks_metrics()); assert!( Operation::Compact { reason: String::new(), rebuild_system_prompt: String::new(), } .tracks_metrics() ); assert!(!Operation::Measure.tracks_metrics()); } }