use crate::stream::StreamEvent;
use crate::tasks::generate::GenerateRequest;
use crate::{InferenceError, InferenceResult};
use serde::{Deserialize, Serialize};
use std::path::PathBuf;
use std::sync::{Arc, OnceLock, RwLock};
#[derive(Clone, Debug, Serialize, Deserialize, PartialEq, Eq)]
pub struct LocalWorkerAdmission {
pub policy: crate::resource_policy::ResourcePolicy,
pub policy_generation: u64,
pub state_root: PathBuf,
pub measured_weights_bytes: u64,
}
#[derive(Clone, Debug, Serialize, Deserialize, PartialEq, Eq)]
pub struct LocalWorkerResidency {
pub model_id: String,
pub measured_weights_bytes: u64,
}
pub struct LocalOffloadResult {
pub result: InferenceResult,
pub residency: LocalWorkerResidency,
pub retention: crate::backend_cache::BackendRetention,
}
pub struct LocalOffloadStream {
pub events: tokio::sync::mpsc::Receiver<StreamEvent>,
pub residency: LocalWorkerResidency,
pub retention: crate::backend_cache::BackendRetention,
}
#[async_trait::async_trait]
pub trait LocalGenerationOffload: Send + Sync {
async fn generate(&self, request: GenerateRequest) -> Result<InferenceResult, InferenceError>;
async fn stream(
&self,
request: GenerateRequest,
) -> Result<tokio::sync::mpsc::Receiver<StreamEvent>, InferenceError>;
async fn generate_admitted(
&self,
request: GenerateRequest,
_admission: LocalWorkerAdmission,
) -> Result<LocalOffloadResult, InferenceError> {
let model_id = request.model.clone().unwrap_or_default();
let result = self.generate(request).await?;
Ok(LocalOffloadResult {
result,
residency: LocalWorkerResidency {
model_id,
measured_weights_bytes: 0,
},
retention: crate::backend_cache::BackendRetention::Transient,
})
}
async fn stream_admitted(
&self,
request: GenerateRequest,
_admission: LocalWorkerAdmission,
) -> Result<LocalOffloadStream, InferenceError> {
let model_id = request.model.clone().unwrap_or_default();
let events = self.stream(request).await?;
Ok(LocalOffloadStream {
events,
residency: LocalWorkerResidency {
model_id,
measured_weights_bytes: 0,
},
retention: crate::backend_cache::BackendRetention::Transient,
})
}
fn refresh_resource_policy(&self, _generation: u64) {}
async fn resident_models(&self) -> Vec<String> {
Vec::new()
}
fn resident_allocation_id(&self, _model_id: &str) -> Option<String> {
None
}
async fn release_model(&self, _model_id: &str) -> Result<bool, InferenceError> {
Ok(false)
}
}
fn offload_slot() -> &'static RwLock<Option<Arc<dyn LocalGenerationOffload>>> {
static SLOT: OnceLock<RwLock<Option<Arc<dyn LocalGenerationOffload>>>> = OnceLock::new();
SLOT.get_or_init(|| RwLock::new(None))
}
pub fn set_local_offload(offload: Option<Arc<dyn LocalGenerationOffload>>) {
let mut guard = offload_slot().write().expect("local offload slot poisoned");
*guard = offload;
}
pub fn current_local_offload() -> Option<Arc<dyn LocalGenerationOffload>> {
if is_offload_worker() {
return None;
}
offload_slot()
.read()
.expect("local offload slot poisoned")
.clone()
}
pub fn is_offload_worker() -> bool {
std::env::var_os("CAR_INFERENCE_WORKER").is_some()
}
#[cfg(test)]
pub(crate) fn test_offload_lock() -> &'static tokio::sync::Mutex<()> {
static LOCK: OnceLock<tokio::sync::Mutex<()>> = OnceLock::new();
LOCK.get_or_init(|| tokio::sync::Mutex::new(()))
}
#[cfg(test)]
mod tests {
use super::*;
struct NoopOffload;
#[async_trait::async_trait]
impl LocalGenerationOffload for NoopOffload {
async fn generate(
&self,
_request: GenerateRequest,
) -> Result<InferenceResult, InferenceError> {
Err(InferenceError::InferenceFailed("noop".into()))
}
async fn stream(
&self,
_request: GenerateRequest,
) -> Result<tokio::sync::mpsc::Receiver<StreamEvent>, InferenceError> {
Err(InferenceError::InferenceFailed("noop".into()))
}
}
#[allow(dead_code)]
fn external_legacy_trait_contract_still_compiles(
implementation: Arc<dyn LocalGenerationOffload>,
) -> Arc<dyn LocalGenerationOffload> {
implementation
}
#[tokio::test]
async fn slot_set_and_clear() {
let _guard = test_offload_lock().lock().await;
assert!(current_local_offload().is_none());
set_local_offload(Some(Arc::new(NoopOffload)));
assert!(current_local_offload().is_some());
set_local_offload(None);
assert!(current_local_offload().is_none());
}
}