use crate::caps::Capabilities;
use crate::handles::Handles;
use crate::infer::{InferBackend, InferResult};
use cuttlefish_abi::{error_codes, Command, Envelope, Event, JobError, JobStatus, Usage};
use std::sync::atomic::{AtomicBool, Ordering};
use std::sync::Arc;
use std::time::Instant;
use tokio::sync::mpsc;
use tokio_util::sync::CancellationToken;
use wasmtime::{Engine, Instance, Linker, Memory, Module, Store, TypedFunc};
#[derive(Debug, Clone)]
pub enum JobEvent {
Token(String),
Progress(serde_json::Value),
}
pub struct JobSpec {
pub module_bytes: Vec<u8>,
pub input: serde_json::Value,
pub caps: Capabilities,
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum Abi {
W32,
W64,
}
impl Abi {
fn ptr_size(self) -> usize {
match self {
Abi::W32 => 4,
Abi::W64 => 8,
}
}
}
struct Guest {
store: Store<()>,
memory: Memory,
abi: Abi,
alloc: TypedFunc<u32, u32>,
init: TypedFunc<(u32, u32), u32>,
step: TypedFunc<(u32, u32), u32>,
on_token: Option<TypedFunc<(u32, u32), i32>>,
}
impl Guest {
fn new(engine: &Engine, module_bytes: &[u8]) -> anyhow::Result<Self> {
let module = Module::new(engine, module_bytes)?;
let linker: Linker<()> = Linker::new(engine);
let mut store = Store::new(engine, ());
let instance: Instance = linker.instantiate(&mut store, &module)?;
let memory = instance
.get_memory(&mut store, "memory")
.ok_or_else(|| anyhow::anyhow!("guest exports no memory"))?;
let abi = if memory.ty(&store).is_64() {
Abi::W64
} else {
Abi::W32
};
if abi == Abi::W64 {
anyhow::bail!("guest uses 64-bit memory; only 32-bit guests are supported");
}
Ok(Self {
alloc: instance.get_typed_func(&mut store, "cf_alloc")?,
init: instance.get_typed_func(&mut store, "cf_init")?,
step: instance.get_typed_func(&mut store, "cf_step")?,
on_token: instance.get_typed_func(&mut store, "cf_on_token").ok(),
memory,
abi,
store,
})
}
fn write(&mut self, bytes: &[u8]) -> anyhow::Result<(u32, u32)> {
let len = bytes.len() as u32;
let ptr = self.alloc.call(&mut self.store, len)?;
self.memory.write(&mut self.store, ptr as usize, bytes)?;
Ok((ptr, len))
}
fn read_desc(&mut self, desc_ptr: u32) -> anyhow::Result<Vec<u8>> {
let w = self.abi.ptr_size();
let mut desc = vec![0u8; 2 * w];
self.memory
.read(&mut self.store, desc_ptr as usize, &mut desc)?;
let field = |bytes: &[u8]| -> u64 {
match w {
4 => u32::from_le_bytes(bytes.try_into().expect("4 bytes")) as u64,
_ => u64::from_le_bytes(bytes.try_into().expect("8 bytes")),
}
};
let ptr = field(&desc[..w]) as usize;
let len = field(&desc[w..]) as usize;
let mut buf = vec![0u8; len];
self.memory.read(&mut self.store, ptr, &mut buf)?;
Ok(buf)
}
fn call_init(&mut self, input: &serde_json::Value) -> anyhow::Result<Command> {
let bytes = serde_json::to_vec(input)?;
let (ptr, len) = self.write(&bytes)?;
let desc = self.init.call(&mut self.store, (ptr, len))?;
Ok(serde_json::from_slice(&self.read_desc(desc)?)?)
}
fn call_step(&mut self, event: &Event) -> anyhow::Result<Command> {
let bytes = serde_json::to_vec(event)?;
let (ptr, len) = self.write(&bytes)?;
let desc = self.step.call(&mut self.store, (ptr, len))?;
Ok(serde_json::from_slice(&self.read_desc(desc)?)?)
}
fn call_on_token(&mut self, token: &str) -> anyhow::Result<bool> {
let Some(f) = self.on_token.clone() else {
return Ok(true);
};
let (ptr, len) = self.write(token.as_bytes())?;
Ok(f.call(&mut self.store, (ptr, len))? == 0)
}
}
fn fail(code: &str, message: impl Into<String>, usage: Usage) -> Envelope {
Envelope {
status: JobStatus::Failed,
result: None,
error: Some(JobError {
code: code.into(),
message: message.into(),
}),
usage,
}
}
fn cancelled(usage: Usage, message: &str) -> Envelope {
Envelope {
status: JobStatus::Cancelled,
result: None,
error: Some(JobError {
code: error_codes::CANCELLED.into(),
message: message.into(),
}),
usage,
}
}
pub async fn run_job(
engine: Arc<Engine>,
backend: Arc<dyn InferBackend>,
job: JobSpec,
events: mpsc::Sender<JobEvent>,
cancel: CancellationToken,
) -> Envelope {
let started = Instant::now();
let mut usage = Usage {
model: backend.model_name(),
..Usage::default()
};
let mut handles = Handles::default();
let mut guest = match Guest::new(&engine, &job.module_bytes) {
Ok(g) => g,
Err(e) => return fail(error_codes::WASM_TRAP, e.to_string(), usage),
};
let mut command = match guest.call_init(&job.input) {
Ok(c) => c,
Err(e) => return fail(error_codes::WASM_TRAP, e.to_string(), usage),
};
loop {
if cancel.is_cancelled() {
usage.duration_ms = started.elapsed().as_millis() as u64;
return cancelled(usage, "job cancelled");
}
let event = match command {
Command::Done { result } => {
usage.duration_ms = started.elapsed().as_millis() as u64;
return Envelope {
status: JobStatus::Completed,
result: Some(result),
error: None,
usage,
};
}
Command::Fail { code, message } => {
usage.duration_ms = started.elapsed().as_millis() as u64;
return fail(&code, message, usage);
}
Command::Emit { progress } => {
let _ = events.send(JobEvent::Progress(progress)).await;
Event::Emitted
}
Command::Open { path } => {
let p = std::path::PathBuf::from(&path);
if !job.caps.allows_read(&p) {
usage.duration_ms = started.elapsed().as_millis() as u64;
return fail(
error_codes::CAPABILITY_DENIED,
format!("read not permitted: {path}"),
usage,
);
}
match handles.open(&p) {
Ok((handle, len)) => Event::Opened { handle, len },
Err(e) => {
usage.duration_ms = started.elapsed().as_millis() as u64;
return fail(error_codes::CAPABILITY_DENIED, e.to_string(), usage);
}
}
}
Command::Slice {
handle,
offset,
len,
} => match handles.slice(handle, offset, len) {
Ok(w) => Event::Sliced {
text: w.text,
next_offset: w.next_offset,
},
Err(e) => {
usage.duration_ms = started.elapsed().as_millis() as u64;
return fail(error_codes::CAPABILITY_DENIED, e.to_string(), usage);
}
},
Command::Infer { prompt, max_tokens } => {
let (tx, mut rx) = mpsc::unbounded_channel::<String>();
let stop = Arc::new(AtomicBool::new(false));
let sink_stop = stop.clone();
let mut sink = move |t: &str| {
tx.send(t.to_string()).is_ok() && !sink_stop.load(Ordering::Relaxed)
};
let mut trap: Option<String> = None;
let outcome: Option<anyhow::Result<InferResult>> = {
let infer = backend.infer(&prompt, max_tokens, &mut sink);
tokio::pin!(infer);
loop {
tokio::select! {
biased;
_ = cancel.cancelled() => break None,
Some(tok) = rx.recv() => {
let _ = events.send(JobEvent::Token(tok.clone())).await;
match guest.call_on_token(&tok) {
Ok(true) => {}
Ok(false) => stop.store(true, Ordering::Relaxed),
Err(e) => {
trap = Some(e.to_string());
break None;
}
}
}
r = &mut infer => break Some(r),
}
}
};
if let Some(message) = trap {
usage.duration_ms = started.elapsed().as_millis() as u64;
return fail(error_codes::WASM_TRAP, message, usage);
}
while let Ok(tok) = rx.try_recv() {
let _ = events.send(JobEvent::Token(tok)).await;
}
match outcome {
None => {
usage.duration_ms = started.elapsed().as_millis() as u64;
return cancelled(usage, "cancelled during inference");
}
Some(Err(e)) => {
usage.duration_ms = started.elapsed().as_millis() as u64;
return fail(error_codes::MODEL_LOAD_FAILED, e.to_string(), usage);
}
Some(Ok(r)) => {
usage.tokens_in += r.tokens_in;
usage.tokens_out += r.tokens_out;
Event::InferDone {
text: r.text,
tokens_out: r.tokens_out,
}
}
}
}
};
command = match guest.call_step(&event) {
Ok(c) => c,
Err(e) => {
usage.duration_ms = started.elapsed().as_millis() as u64;
return fail(error_codes::WASM_TRAP, e.to_string(), usage);
}
};
}
}