use crate::stream::StreamEvent;
use crate::tasks::generate::GenerateRequest;
use crate::{InferenceError, InferenceResult};
use serde::{Deserialize, Serialize};
use std::future::Future;
use std::path::PathBuf;
use std::sync::atomic::{AtomicBool, Ordering};
use std::sync::{Arc, OnceLock, RwLock};
#[derive(Clone, Debug)]
struct InferenceControlScope {
inference_id: String,
termination: ControlledTerminationToken,
}
tokio::task_local! {
static INFERENCE_CONTROL_SCOPE: InferenceControlScope;
}
#[derive(Clone, Debug, Default)]
pub struct ControlledTerminationToken(Arc<AtomicBool>);
impl ControlledTerminationToken {
pub fn confirm_exact_backend_termination(&self) {
self.0.store(true, Ordering::Release);
}
pub fn is_confirmed(&self) -> bool {
self.0.load(Ordering::Acquire)
}
pub fn error_if_confirmed(&self) -> Result<(), InferenceError> {
if self.is_confirmed() {
Err(InferenceError::ControlledTermination)
} else {
Ok(())
}
}
}
pub async fn scope_inference_control_id<F>(inference_id: String, future: F) -> F::Output
where
F: Future,
{
INFERENCE_CONTROL_SCOPE
.scope(
InferenceControlScope {
inference_id,
termination: ControlledTerminationToken::default(),
},
future,
)
.await
}
pub fn current_inference_control_id() -> Option<String> {
INFERENCE_CONTROL_SCOPE
.try_with(|scope| scope.inference_id.clone())
.ok()
}
pub fn current_controlled_termination_token() -> Option<ControlledTerminationToken> {
INFERENCE_CONTROL_SCOPE
.try_with(|scope| scope.termination.clone())
.ok()
}
pub fn ensure_not_controlled_terminated() -> Result<(), InferenceError> {
match current_controlled_termination_token() {
Some(token) => token.error_if_confirmed(),
None => Ok(()),
}
}
#[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,
}
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
pub enum InferenceTerminationAck {
Confirmed,
Unconfirmed,
}
#[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)
}
async fn terminate_inference(&self, _inference_id: &str) -> InferenceTerminationAck {
InferenceTerminationAck::Unconfirmed
}
}
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;
#[tokio::test]
async fn legacy_or_nonisolated_offload_never_confirms_termination() {
let offload = NoopOffload;
assert_eq!(
offload.terminate_inference("opaque").await,
InferenceTerminationAck::Unconfirmed
);
}
#[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());
}
}