use std::num::NonZeroUsize;
use std::sync::Arc;
use std::sync::atomic::{AtomicU32, Ordering};
use std::time::Duration;
use async_trait::async_trait;
use iroh::EndpointId as NodeId;
use iroh::endpoint::{Connection, Endpoint, RecvStream, SendStream};
use iroh::protocol::{AcceptError, ProtocolHandler};
use iroh_blobs::Hash;
use serde::{Deserialize, Serialize};
use tracing::{debug, warn};
use uuid::Uuid;
use super::runtime::{CompiledTask, ExecMetrics, HostGrants, TaskError, WasmRuntime};
use super::{COMPUTE_ALPN, ResourceLimits, TaskClass};
pub const MAX_REQUEST_BYTES: usize = 16 * 1024 * 1024;
pub const MAX_REPLY_BYTES: usize = 16 * 1024 * 1024;
const COMPILED_CACHE_ENTRIES: usize = 32;
#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)]
pub enum ComputeRequest {
Execute(ExecuteRequest),
Probe(ProbeRequest),
}
#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)]
pub struct ProbeRequest {
pub class: TaskClass,
}
#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)]
pub struct ProbeReply {
pub accepts_class: bool,
pub free_slots: u32,
pub cpu_load_pct: u8,
pub ram_free_mb: u32,
}
#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)]
pub struct ExecuteRequest {
pub task_id: Uuid,
pub wasm_hash: Hash,
pub entrypoint: String,
pub class: TaskClass,
pub limits: ResourceLimits,
pub input: Vec<u8>,
pub required_model: Option<String>,
}
#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)]
pub enum ExecuteAck {
Accepted,
Rejected(RejectReason),
}
#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize, thiserror::Error)]
pub enum RejectReason {
#[error("executor does not accept tasks of this class")]
ClassNotAccepted,
#[error("executor has no free concurrency slot")]
Busy,
#[error("executor does not serve the required NN model: {0}")]
ModelNotAvailable(String),
#[error("malformed request: {0}")]
Malformed(String),
}
#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)]
pub struct ExecuteReply {
pub outcome: Result<CompletedTask, TaskError>,
}
#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)]
pub struct CompletedTask {
pub output: Vec<u8>,
pub metrics: ExecMetrics,
}
#[derive(Debug, Clone)]
pub struct ExecutorPolicy {
pub accepts: Vec<TaskClass>,
pub max_concurrent: u32,
}
impl Default for ExecutorPolicy {
fn default() -> Self {
Self {
accepts: vec![TaskClass::General, TaskClass::Media, TaskClass::Analytics],
max_concurrent: 2,
}
}
}
#[async_trait]
pub trait WasmFetcher: Send + Sync + 'static {
async fn fetch_wasm(&self, hash: &Hash, provider: NodeId) -> Result<Vec<u8>, String>;
}
#[async_trait]
impl WasmFetcher for crate::p2p::network::core::blobs::BlobStore {
async fn fetch_wasm(&self, hash: &Hash, provider: NodeId) -> Result<Vec<u8>, String> {
self.get_or_download(hash, &[provider])
.await
.map(|bytes| bytes.to_vec())
.map_err(|e| e.to_string())
}
}
fn admit(
policy: &ExecutorPolicy,
class: TaskClass,
running: &Arc<AtomicU32>,
) -> Result<SlotGuard, RejectReason> {
if !policy.accepts.contains(&class) {
return Err(RejectReason::ClassNotAccepted);
}
let prev = running.fetch_add(1, Ordering::AcqRel);
if prev >= policy.max_concurrent {
running.fetch_sub(1, Ordering::AcqRel);
return Err(RejectReason::Busy);
}
Ok(SlotGuard {
running: running.clone(),
})
}
struct SlotGuard {
running: Arc<AtomicU32>,
}
impl Drop for SlotGuard {
fn drop(&mut self) {
self.running.fetch_sub(1, Ordering::AcqRel);
}
}
#[derive(Clone)]
pub struct ComputeProtocolHandler {
runtime: Arc<WasmRuntime>,
fetcher: Arc<dyn WasmFetcher>,
policy: Arc<parking_lot::RwLock<ExecutorPolicy>>,
grants: Arc<parking_lot::RwLock<HostGrants>>,
#[cfg(feature = "compute-nn")]
nn_models: Arc<parking_lot::RwLock<Option<Arc<crate::compute::nn::NnModelRegistry>>>>,
running: Arc<AtomicU32>,
compiled: Arc<parking_lot::Mutex<lru::LruCache<Hash, CompiledTask>>>,
probe_system: Arc<parking_lot::Mutex<sysinfo::System>>,
}
impl std::fmt::Debug for ComputeProtocolHandler {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct("ComputeProtocolHandler")
.field("running", &self.running.load(Ordering::Relaxed))
.finish_non_exhaustive()
}
}
impl ComputeProtocolHandler {
pub fn new(fetcher: Arc<dyn WasmFetcher>, policy: ExecutorPolicy) -> Result<Self, TaskError> {
Ok(Self {
runtime: Arc::new(WasmRuntime::new()?),
fetcher,
policy: Arc::new(parking_lot::RwLock::new(policy)),
grants: Arc::new(parking_lot::RwLock::new(HostGrants::default())),
#[cfg(feature = "compute-nn")]
nn_models: Arc::new(parking_lot::RwLock::new(None)),
running: Arc::new(AtomicU32::new(0)),
compiled: Arc::new(parking_lot::Mutex::new(lru::LruCache::new(
NonZeroUsize::new(COMPILED_CACHE_ENTRIES).expect("nonzero"),
))),
probe_system: Arc::new(parking_lot::Mutex::new({
let mut system = sysinfo::System::new();
system.refresh_cpu_usage();
system.refresh_memory();
system
})),
})
}
pub fn set_policy(&self, policy: ExecutorPolicy) {
*self.policy.write() = policy;
}
pub fn set_host_grants(&self, grants: HostGrants) {
*self.grants.write() = grants;
}
#[cfg(feature = "compute-nn")]
pub fn set_nn_models(&self, registry: Arc<crate::compute::nn::NnModelRegistry>) {
*self.nn_models.write() = Some(registry);
let mut policy = self.policy.write();
if !policy.accepts.contains(&TaskClass::Inference) {
policy.accepts.push(TaskClass::Inference);
}
}
async fn probe_reply(&self, class: TaskClass) -> ProbeReply {
let policy = self.policy.read().clone();
let running = self.running.load(Ordering::Relaxed);
let system = self.probe_system.clone();
let (cpu_load_pct, ram_free_mb) = tokio::task::spawn_blocking(move || {
let mut system = system.lock();
system.refresh_cpu_usage();
system.refresh_memory();
(
system.global_cpu_usage().round().clamp(0.0, 100.0) as u8,
(system.available_memory() / (1024 * 1024)).min(u32::MAX as u64) as u32,
)
})
.await
.unwrap_or((0, 0));
ProbeReply {
accepts_class: policy.accepts.contains(&class) && running < policy.max_concurrent,
free_slots: policy.max_concurrent.saturating_sub(running),
cpu_load_pct,
ram_free_mb,
}
}
pub fn policy(&self) -> ExecutorPolicy {
self.policy.read().clone()
}
pub fn tasks_running(&self) -> u32 {
self.running.load(Ordering::Relaxed)
}
pub fn nn_model_names(&self) -> Vec<String> {
#[cfg(feature = "compute-nn")]
{
self.nn_models
.read()
.as_ref()
.map(|registry| registry.model_names())
.unwrap_or_default()
}
#[cfg(not(feature = "compute-nn"))]
{
Vec::new()
}
}
fn serves_model(&self, name: &str) -> bool {
#[cfg(feature = "compute-nn")]
{
let in_registry = self
.nn_models
.read()
.as_ref()
.is_some_and(|registry| registry.has_model(name));
let in_grants = self
.grants
.read()
.nn
.as_ref()
.is_some_and(|grant| grant.has_model(name));
in_registry || in_grants
}
#[cfg(not(feature = "compute-nn"))]
{
let _ = name;
false
}
}
async fn compiled_task(
&self,
hash: &Hash,
provider: NodeId,
) -> Result<CompiledTask, TaskError> {
if let Some(task) = self.compiled.lock().get(hash).cloned() {
debug!(hash = %hash.fmt_short(), "compute: compiled-module cache hit");
return Ok(task);
}
let wasm = self
.fetcher
.fetch_wasm(hash, provider)
.await
.map_err(TaskError::WasmUnavailable)?;
let task = self.runtime.compile(&wasm)?;
self.compiled.lock().put(*hash, task.clone());
Ok(task)
}
async fn serve(
&self,
requester: NodeId,
send: &mut SendStream,
recv: &mut RecvStream,
) -> Result<(), AcceptError> {
let raw = read_frame(recv, MAX_REQUEST_BYTES)
.await
.map_err(AcceptError::from_err)?;
let request: ComputeRequest = match postcard::from_bytes(&raw) {
Ok(req) => req,
Err(e) => {
write_frame(
send,
&encode(&ExecuteAck::Rejected(RejectReason::Malformed(
e.to_string(),
)))?,
)
.await
.map_err(AcceptError::from_err)?;
return Ok(());
}
};
let request = match request {
ComputeRequest::Execute(request) => request,
ComputeRequest::Probe(probe) => {
let reply = self.probe_reply(probe.class).await;
write_frame(send, &encode(&reply)?)
.await
.map_err(AcceptError::from_err)?;
return Ok(());
}
};
if let Some(model) = &request.required_model
&& !self.serves_model(model)
{
write_frame(
send,
&encode(&ExecuteAck::Rejected(RejectReason::ModelNotAvailable(
model.clone(),
)))?,
)
.await
.map_err(AcceptError::from_err)?;
return Ok(());
}
let slot = {
let policy = self.policy.read().clone();
admit(&policy, request.class, &self.running)
};
let slot = match slot {
Ok(slot) => slot,
Err(reason) => {
debug!(task = %request.task_id, peer = %requester.fmt_short(),
%reason, "compute: task rejected");
write_frame(send, &encode(&ExecuteAck::Rejected(reason))?)
.await
.map_err(AcceptError::from_err)?;
return Ok(());
}
};
write_frame(send, &encode(&ExecuteAck::Accepted)?)
.await
.map_err(AcceptError::from_err)?;
let outcome = match self.compiled_task(&request.wasm_hash, requester).await {
Ok(task) => {
let runtime = self.runtime.clone();
let entrypoint = request.entrypoint.clone();
let limits = request.limits;
let input = request.input;
#[allow(unused_mut)]
let mut grants = self.grants.read().clone();
#[cfg(feature = "compute-nn")]
let nn_registry = self.nn_models.read().clone();
#[cfg(feature = "compute-nn")]
if (request.class == TaskClass::Inference || request.required_model.is_some())
&& grants.nn.is_none()
&& let Some(registry) = nn_registry
&& !registry.is_empty()
{
match registry.grant_for(&self.fetcher, requester).await {
Ok(grant) => grants.nn = Some(grant),
Err(e) => {
warn!(task = %request.task_id, error = %e,
"compute: NN model preparation failed");
write_frame(send, &encode(&ExecuteReply { outcome: Err(e) })?)
.await
.map_err(AcceptError::from_err)?;
return Ok(());
}
}
}
let response_deadline = Duration::from_millis(
limits.timeout_ms.saturating_mul(2).saturating_add(1_000),
);
let run = tokio::task::spawn_blocking(move || {
runtime.execute_with_host(&task, &entrypoint, &input, &limits, &grants)
});
match tokio::time::timeout(response_deadline, run).await {
Err(_elapsed) => Err(TaskError::DeadlineExceeded),
Ok(join) => join
.map_err(|e| TaskError::Runtime(format!("executor task panicked: {e}")))
.and_then(|r| r)
.map(|exec| CompletedTask {
output: exec.output,
metrics: exec.metrics,
}),
}
}
Err(e) => Err(e),
};
drop(slot);
match &outcome {
Ok(done) => debug!(task = %request.task_id, peer = %requester.fmt_short(),
fuel = done.metrics.fuel_consumed, ms = done.metrics.duration_ms,
"compute: task completed"),
Err(e) => warn!(task = %request.task_id, peer = %requester.fmt_short(),
error = %e, "compute: task failed"),
}
write_frame(send, &encode(&ExecuteReply { outcome })?)
.await
.map_err(AcceptError::from_err)?;
Ok(())
}
}
impl ProtocolHandler for ComputeProtocolHandler {
async fn accept(&self, connection: Connection) -> Result<(), AcceptError> {
let requester = connection.remote_id();
let (mut send, mut recv) = connection.accept_bi().await?;
self.serve(requester, &mut send, &mut recv).await?;
send.finish().map_err(AcceptError::from_err)?;
connection.closed().await;
Ok(())
}
}
#[derive(Debug, Clone, PartialEq, Eq, thiserror::Error)]
pub enum ComputeCallError {
#[error("executor unreachable: {0}")]
Unreachable(String),
#[error("executor rejected the task: {0}")]
Rejected(RejectReason),
#[error("task failed on the executor: {0}")]
Task(TaskError),
#[error("protocol error: {0}")]
Protocol(String),
#[error("timed out waiting for the executor")]
Timeout,
}
#[derive(Debug, Clone)]
pub struct ComputeClient {
endpoint: Endpoint,
}
impl ComputeClient {
pub fn new(endpoint: Endpoint) -> Self {
Self { endpoint }
}
pub async fn execute_on(
&self,
executor: impl Into<iroh::EndpointAddr>,
request: ExecuteRequest,
total_timeout: Duration,
) -> Result<CompletedTask, ComputeCallError> {
tokio::time::timeout(total_timeout, self.call(executor.into(), request))
.await
.map_err(|_| ComputeCallError::Timeout)?
}
pub async fn probe(
&self,
executor: impl Into<iroh::EndpointAddr>,
class: TaskClass,
timeout: Duration,
) -> Result<ProbeReply, ComputeCallError> {
let executor = executor.into();
tokio::time::timeout(timeout, async move {
let connection = self
.endpoint
.connect(executor, COMPUTE_ALPN)
.await
.map_err(|e| ComputeCallError::Unreachable(e.to_string()))?;
let (mut send, mut recv) = connection
.open_bi()
.await
.map_err(|e| ComputeCallError::Unreachable(e.to_string()))?;
let raw = encode(&ComputeRequest::Probe(ProbeRequest { class }))
.map_err(|e| ComputeCallError::Protocol(format!("encode: {e}")))?;
write_frame(&mut send, &raw)
.await
.map_err(|e| ComputeCallError::Protocol(format!("send probe: {e}")))?;
send.finish()
.map_err(|e| ComputeCallError::Protocol(format!("finish stream: {e}")))?;
let reply_raw = read_frame(&mut recv, 4096)
.await
.map_err(|e| ComputeCallError::Protocol(format!("read probe reply: {e}")))?;
let reply: ProbeReply = postcard::from_bytes(&reply_raw)
.map_err(|e| ComputeCallError::Protocol(format!("decode probe reply: {e}")))?;
connection.close(0u32.into(), b"done");
Ok(reply)
})
.await
.map_err(|_| ComputeCallError::Timeout)?
}
async fn call(
&self,
executor: iroh::EndpointAddr,
request: ExecuteRequest,
) -> Result<CompletedTask, ComputeCallError> {
let connection = self
.endpoint
.connect(executor, COMPUTE_ALPN)
.await
.map_err(|e| ComputeCallError::Unreachable(e.to_string()))?;
let (mut send, mut recv) = connection
.open_bi()
.await
.map_err(|e| ComputeCallError::Unreachable(e.to_string()))?;
let raw = encode(&ComputeRequest::Execute(request))
.map_err(|e| ComputeCallError::Protocol(format!("encode: {e}")))?;
write_frame(&mut send, &raw)
.await
.map_err(|e| ComputeCallError::Protocol(format!("send request: {e}")))?;
send.finish()
.map_err(|e| ComputeCallError::Protocol(format!("finish stream: {e}")))?;
let ack_raw = read_frame(&mut recv, 4096)
.await
.map_err(|e| ComputeCallError::Protocol(format!("read ack: {e}")))?;
let ack: ExecuteAck = postcard::from_bytes(&ack_raw)
.map_err(|e| ComputeCallError::Protocol(format!("decode ack: {e}")))?;
if let ExecuteAck::Rejected(reason) = ack {
connection.close(0u32.into(), b"rejected");
return Err(ComputeCallError::Rejected(reason));
}
let reply_raw = read_frame(&mut recv, MAX_REPLY_BYTES)
.await
.map_err(|e| ComputeCallError::Protocol(format!("read reply: {e}")))?;
let reply: ExecuteReply = postcard::from_bytes(&reply_raw)
.map_err(|e| ComputeCallError::Protocol(format!("decode reply: {e}")))?;
connection.close(0u32.into(), b"done");
reply.outcome.map_err(ComputeCallError::Task)
}
}
fn encode<T: Serialize>(msg: &T) -> Result<Vec<u8>, AcceptError> {
postcard::to_stdvec(msg).map_err(AcceptError::from_err)
}
async fn write_frame(send: &mut SendStream, payload: &[u8]) -> std::io::Result<()> {
let len = u32::try_from(payload.len())
.map_err(|_| std::io::Error::new(std::io::ErrorKind::InvalidInput, "frame too large"))?;
send.write_all(&len.to_le_bytes())
.await
.map_err(std::io::Error::other)?;
send.write_all(payload)
.await
.map_err(std::io::Error::other)?;
Ok(())
}
async fn read_frame(recv: &mut RecvStream, max: usize) -> std::io::Result<Vec<u8>> {
let mut len_buf = [0u8; 4];
recv.read_exact(&mut len_buf)
.await
.map_err(std::io::Error::other)?;
let len = u32::from_le_bytes(len_buf) as usize;
if len > max {
return Err(std::io::Error::new(
std::io::ErrorKind::InvalidData,
format!("frame of {len} bytes exceeds the {max}-byte limit"),
));
}
let mut payload = vec![0u8; len];
recv.read_exact(&mut payload)
.await
.map_err(std::io::Error::other)?;
Ok(payload)
}
#[cfg(test)]
mod tests {
use super::*;
fn sample_request() -> ExecuteRequest {
ExecuteRequest {
task_id: Uuid::new_v4(),
wasm_hash: Hash::new(b"some wasm"),
entrypoint: "gdb_run".into(),
class: TaskClass::General,
limits: ResourceLimits::default(),
input: b"payload".to_vec(),
required_model: None,
}
}
#[test]
fn request_roundtrips_through_postcard() {
let req = ComputeRequest::Execute(sample_request());
let bytes = postcard::to_stdvec(&req).expect("serialize");
let back: ComputeRequest = postcard::from_bytes(&bytes).expect("deserialize");
assert_eq!(req, back);
let probe = ComputeRequest::Probe(ProbeRequest {
class: TaskClass::Analytics,
});
let bytes = postcard::to_stdvec(&probe).expect("serialize");
let back: ComputeRequest = postcard::from_bytes(&bytes).expect("deserialize");
assert_eq!(probe, back);
}
#[tokio::test]
async fn probe_reply_reflects_policy_and_slots() {
struct NoFetch;
#[async_trait]
impl WasmFetcher for NoFetch {
async fn fetch_wasm(&self, _h: &Hash, _p: NodeId) -> Result<Vec<u8>, String> {
Err("unused".into())
}
}
let handler = ComputeProtocolHandler::new(Arc::new(NoFetch), ExecutorPolicy::default())
.expect("handler");
let reply = handler.probe_reply(TaskClass::General).await;
assert!(reply.accepts_class);
assert_eq!(reply.free_slots, 2);
assert!(reply.cpu_load_pct <= 100);
let reply = handler.probe_reply(TaskClass::Inference).await;
assert!(!reply.accepts_class);
handler.set_policy(ExecutorPolicy {
max_concurrent: 0,
..ExecutorPolicy::default()
});
let reply = handler.probe_reply(TaskClass::General).await;
assert!(!reply.accepts_class);
assert_eq!(reply.free_slots, 0);
}
#[test]
fn reply_roundtrips_including_task_error() {
let reply = ExecuteReply {
outcome: Err(TaskError::FuelExhausted),
};
let bytes = postcard::to_stdvec(&reply).expect("serialize");
let back: ExecuteReply = postcard::from_bytes(&bytes).expect("deserialize");
assert_eq!(reply, back);
}
#[cfg(feature = "compute-nn")]
#[tokio::test]
async fn set_nn_models_enables_inference_admission() {
struct NoFetch;
#[async_trait]
impl WasmFetcher for NoFetch {
async fn fetch_wasm(&self, _h: &Hash, _p: NodeId) -> Result<Vec<u8>, String> {
Err("unused".into())
}
}
let handler = ComputeProtocolHandler::new(Arc::new(NoFetch), ExecutorPolicy::default())
.expect("handler");
assert!(!handler.policy().accepts.contains(&TaskClass::Inference));
handler.set_nn_models(Arc::new(crate::compute::nn::NnModelRegistry::new()));
assert!(handler.policy().accepts.contains(&TaskClass::Inference));
assert!(
handler
.probe_reply(TaskClass::Inference)
.await
.accepts_class
);
}
#[test]
fn admission_rejects_class_not_in_policy() {
let running = Arc::new(AtomicU32::new(0));
let policy = ExecutorPolicy::default();
assert!(matches!(
admit(&policy, TaskClass::Inference, &running),
Err(RejectReason::ClassNotAccepted)
));
assert_eq!(running.load(Ordering::Relaxed), 0);
}
#[test]
fn admission_enforces_concurrency_and_releases_slots() {
let running = Arc::new(AtomicU32::new(0));
let policy = ExecutorPolicy {
accepts: vec![TaskClass::General],
max_concurrent: 2,
};
let a = admit(&policy, TaskClass::General, &running).expect("slot 1");
let _b = admit(&policy, TaskClass::General, &running).expect("slot 2");
assert!(matches!(
admit(&policy, TaskClass::General, &running),
Err(RejectReason::Busy)
));
assert_eq!(running.load(Ordering::Relaxed), 2);
drop(a);
assert_eq!(running.load(Ordering::Relaxed), 1);
let _c = admit(&policy, TaskClass::General, &running).expect("slot freed");
}
#[test]
fn zero_concurrency_disables_execution() {
let running = Arc::new(AtomicU32::new(0));
let policy = ExecutorPolicy {
accepts: vec![TaskClass::General],
max_concurrent: 0,
};
assert!(matches!(
admit(&policy, TaskClass::General, &running),
Err(RejectReason::Busy)
));
assert_eq!(running.load(Ordering::Relaxed), 0);
}
}