use crate::caps::Capabilities;
use crate::handles::Handles;
use crate::infer::{InferBackend, InferRequest, 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};
const RENDER_WIDTH: u16 = 1024;
#[derive(Debug, Clone)]
pub enum JobEvent {
Token(String),
Progress(serde_json::Value),
}
pub struct JobSpec {
pub nodes: Vec<crate::dag::CheckedNode>,
pub exclusive_to: std::collections::HashMap<String, crate::dag::BranchExclusivity>,
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)
}
}
pub fn read_signature(
engine: &Engine,
module_bytes: &[u8],
) -> anyhow::Result<cuttlefish_abi::Signature> {
let permissive = cuttlefish_abi::Signature {
input: cuttlefish_abi::Ty::Json,
output: cuttlefish_abi::Ty::Json,
};
let module = Module::new(engine, module_bytes)?;
let linker: Linker<()> = Linker::new(engine);
let mut store = Store::new(engine, ());
let instance = linker.instantiate(&mut store, &module)?;
let Ok(signature) = instance.get_typed_func::<(), u32>(&mut store, "cf_signature") else {
return Ok(permissive);
};
let Some(memory) = instance.get_memory(&mut store, "memory") else {
return Ok(permissive);
};
let desc_ptr = signature.call(&mut store, ())? as usize;
let mut desc = [0u8; 8];
memory.read(&mut store, desc_ptr, &mut desc)?;
let ptr = u32::from_le_bytes(desc[..4].try_into().expect("4 bytes")) as usize;
let len = u32::from_le_bytes(desc[4..].try_into().expect("4 bytes")) as usize;
let mut buf = vec![0u8; len];
memory.read(&mut store, ptr, &mut buf)?;
Ok(serde_json::from_slice(&buf)?)
}
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 references_any(
expr: &cuttlefish_core::graph::InputExpr,
skipped: &std::collections::HashSet<String>,
) -> bool {
use cuttlefish_core::graph::InputExpr;
match expr {
InputExpr::FromNode(n) => skipped.contains(n),
InputExpr::Record(fields) => fields.values().any(|e| references_any(e, skipped)),
InputExpr::List(items) => items.iter().any(|e| references_any(e, skipped)),
}
}
fn evaluate_input(
expr: &cuttlefish_core::graph::InputExpr,
outputs: &std::collections::HashMap<String, serde_json::Value>,
) -> serde_json::Value {
use cuttlefish_core::graph::InputExpr;
match expr {
InputExpr::FromNode(n) => outputs.get(n).cloned().unwrap_or(serde_json::Value::Null),
InputExpr::Record(fields) => {
let mut map = serde_json::Map::new();
for (k, v) in fields {
map.insert(k.clone(), evaluate_input(v, outputs));
}
serde_json::Value::Object(map)
}
InputExpr::List(items) => {
serde_json::Value::Array(items.iter().map(|e| evaluate_input(e, outputs)).collect())
}
}
}
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,
ledger: &crate::ledger::Ledger,
) -> Envelope {
let started = Instant::now();
let mut usage = Usage {
model: backend.model_name(),
..Usage::default()
};
let mut handles = Handles::default();
let envelope = 'run: {
if job.nodes.is_empty() {
usage.duration_ms = started.elapsed().as_millis() as u64;
break 'run fail(
error_codes::SCHEMA_VALIDATION_FAILED,
"this job has no nodes to run",
usage,
);
}
let total = job.nodes.len();
let mut outputs: std::collections::HashMap<String, serde_json::Value> =
std::collections::HashMap::new();
let mut skipped: std::collections::HashSet<String> = std::collections::HashSet::new();
let mut route_taken: std::collections::HashMap<String, String> =
std::collections::HashMap::new();
for (index, node) in job.nodes.iter().enumerate() {
match ledger.is_skipped(&node.name) {
Ok(true) => {
skipped.insert(node.name.clone());
continue;
}
Ok(false) => {}
Err(e) => {
usage.duration_ms = started.elapsed().as_millis() as u64;
break 'run fail(
error_codes::SCHEMA_VALIDATION_FAILED,
format!("reading ledger skip state for node `{}`: {e}", node.name),
usage,
);
}
}
match ledger.get_completed(&node.name) {
Ok(Some(cached)) => {
if let Some(route) = cached.get("route").and_then(|v| v.as_str()) {
route_taken.insert(node.name.clone(), route.to_string());
}
outputs.insert(node.name.clone(), cached);
continue;
}
Ok(None) => {}
Err(e) => {
usage.duration_ms = started.elapsed().as_millis() as u64;
break 'run fail(
error_codes::SCHEMA_VALIDATION_FAILED,
format!("reading ledger checkpoint for node `{}`: {e}", node.name),
usage,
);
}
}
if total > 1 {
let _ = events
.send(JobEvent::Progress(serde_json::json!({
"stage": index + 1,
"of": total,
"node": node.name,
})))
.await;
}
if let Some(ex) = job.exclusive_to.get(&node.name) {
if let Some(taken) = route_taken.get(&ex.decision) {
if taken != &ex.label {
skipped.insert(node.name.clone());
if let Err(e) = ledger.write_skipped(&node.name) {
usage.duration_ms = started.elapsed().as_millis() as u64;
break 'run fail(
error_codes::SCHEMA_VALIDATION_FAILED,
format!("recording skip for node `{}` in ledger: {e}", node.name),
usage,
);
}
continue;
}
}
}
if let Some(expr) = &node.input {
if references_any(expr, &skipped) {
skipped.insert(node.name.clone());
if let Err(e) = ledger.write_skipped(&node.name) {
usage.duration_ms = started.elapsed().as_millis() as u64;
break 'run fail(
error_codes::SCHEMA_VALIDATION_FAILED,
format!("recording skip for node `{}` in ledger: {e}", node.name),
usage,
);
}
continue;
}
}
let node_input = match &node.input {
None => job.input.clone(),
Some(expr) => evaluate_input(expr, &outputs),
};
let mut current_input = node_input;
let mut iterations: u32 = 0;
let result = loop {
let r = match run_stage(
&engine,
&backend,
&node.module_bytes,
current_input.clone(),
&job.caps,
&mut handles,
&events,
&cancel,
&mut usage,
started,
index,
)
.await
{
Ok(v) => v,
Err(envelope) => break 'run envelope,
};
match &node.repeat_until {
None => break r,
Some(field) => {
iterations += 1;
let done = r.get(field).and_then(|v| v.as_str()) == Some("done");
if done {
break r;
}
let max = node
.max_iterations
.expect("repeat_until requires max_iterations, enforced at parse time");
if iterations >= max {
usage.duration_ms = started.elapsed().as_millis() as u64;
break 'run fail(
error_codes::SCHEMA_VALIDATION_FAILED,
format!(
"node `{}` did not reach repeat_until=\"done\" within max_iterations={max}",
node.name
),
usage,
);
}
current_input = r;
}
}
};
let valid_labels: Vec<&str> = job
.exclusive_to
.values()
.filter(|ex| ex.decision == node.name)
.map(|ex| ex.label.as_str())
.collect();
if !valid_labels.is_empty() {
match result.get("route").and_then(|v| v.as_str()) {
Some(route) if valid_labels.contains(&route) => {
route_taken.insert(node.name.clone(), route.to_string());
}
Some(route) => {
usage.duration_ms = started.elapsed().as_millis() as u64;
break 'run fail(
error_codes::SCHEMA_VALIDATION_FAILED,
format!(
"node `{}` produced route \"{route}\", which doesn't match any \
declared label for this branches decision",
node.name
),
usage,
);
}
None => {
usage.duration_ms = started.elapsed().as_millis() as u64;
break 'run fail(
error_codes::SCHEMA_VALIDATION_FAILED,
format!(
"node `{}` is a branches decision but its output has no `route` field",
node.name
),
usage,
);
}
}
} else if let Some(route) = result.get("route").and_then(|v| v.as_str()) {
route_taken.insert(node.name.clone(), route.to_string());
}
if let Err(e) = ledger.write_completed(&node.name, &result) {
usage.duration_ms = started.elapsed().as_millis() as u64;
break 'run fail(
error_codes::SCHEMA_VALIDATION_FAILED,
format!(
"recording checkpoint for node `{}` in ledger: {e}",
node.name
),
usage,
);
}
outputs.insert(node.name.clone(), result);
}
usage.duration_ms = started.elapsed().as_millis() as u64;
let result = job
.nodes
.iter()
.rev()
.find_map(|n| outputs.get(&n.name).cloned());
match result {
Some(value) => Envelope {
status: JobStatus::Completed,
result: Some(value),
error: None,
usage,
},
None => fail(
error_codes::SCHEMA_VALIDATION_FAILED,
"every node in this job was skipped; there is no result",
usage,
),
}
};
let ledger_status = match envelope.status {
JobStatus::Completed => "completed",
JobStatus::Failed => "failed",
JobStatus::Cancelled => "cancelled",
_ => "running", };
if let Err(e) = ledger.finish(ledger_status) {
eprintln!("warning: failed to record job {ledger_status} status in ledger: {e}");
}
envelope
}
#[allow(clippy::too_many_arguments)]
async fn run_stage(
engine: &Engine,
backend: &Arc<dyn InferBackend>,
module_bytes: &[u8],
input: serde_json::Value,
caps: &Capabilities,
handles: &mut Handles,
events: &mpsc::Sender<JobEvent>,
cancel: &CancellationToken,
usage: &mut Usage,
started: Instant,
stage_index: usize,
) -> Result<serde_json::Value, Envelope> {
let blame = |message: String| -> String {
if stage_index == 0 {
message
} else {
format!("block {} of the pipeline: {message}", stage_index + 1)
}
};
let mut doc_paths: std::collections::HashMap<u32, std::path::PathBuf> =
std::collections::HashMap::new();
let mut guest = match Guest::new(engine, module_bytes) {
Ok(g) => g,
Err(e) => {
return Err(fail(
error_codes::WASM_TRAP,
blame(e.to_string()),
usage.clone(),
))
}
};
let mut command = match guest.call_init(&input) {
Ok(c) => c,
Err(e) => {
return Err(fail(
error_codes::WASM_TRAP,
blame(e.to_string()),
usage.clone(),
))
}
};
loop {
if cancel.is_cancelled() {
usage.duration_ms = started.elapsed().as_millis() as u64;
return Err(cancelled(usage.clone(), "job cancelled"));
}
let event = match command {
Command::Done { result } => return Ok(result),
Command::Fail { code, message } => {
usage.duration_ms = started.elapsed().as_millis() as u64;
return Err(fail(&code, blame(message), usage.clone()));
}
Command::Emit { progress } => {
let _ = events.send(JobEvent::Progress(progress)).await;
Event::Emitted
}
Command::Open { path } => {
let p = std::path::PathBuf::from(&path);
if !caps.allows_read(&p) {
usage.duration_ms = started.elapsed().as_millis() as u64;
return Err(fail(
error_codes::CAPABILITY_DENIED,
format!("read not permitted: {path}"),
usage.clone(),
));
}
match handles.open(&p) {
Ok((handle, len, kind)) => {
let kind = match kind {
cuttlefish_abi::MediaKind::Document { .. } => {
match crate::documents::inspect(&p) {
Ok(info) => cuttlefish_abi::MediaKind::Document {
pages: info.pages,
has_text_layer: info.has_text_layer,
},
Err(_) => cuttlefish_abi::MediaKind::Binary,
}
}
other => other,
};
doc_paths.insert(handle, p.clone());
Event::Opened { handle, len, kind }
}
Err(e) => {
usage.duration_ms = started.elapsed().as_millis() as u64;
return Err(fail(
error_codes::CAPABILITY_DENIED,
e.to_string(),
usage.clone(),
));
}
}
}
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 Err(fail(
error_codes::CAPABILITY_DENIED,
e.to_string(),
usage.clone(),
));
}
},
Command::SliceBytes {
handle,
offset,
len,
} => match handles.slice_bytes(handle, offset, len) {
Ok((bytes, next_offset)) => {
use base64::Engine;
Event::SlicedBytes {
bytes_base64: base64::engine::general_purpose::STANDARD.encode(&bytes),
next_offset,
}
}
Err(e) => {
usage.duration_ms = started.elapsed().as_millis() as u64;
return Err(fail(
error_codes::CAPABILITY_DENIED,
e.to_string(),
usage.clone(),
));
}
},
Command::PageText { handle, page } => {
let Some(path) = doc_paths.get(&handle).cloned() else {
usage.duration_ms = started.elapsed().as_millis() as u64;
return Err(fail(
error_codes::CAPABILITY_DENIED,
format!("no such handle: {handle}"),
usage.clone(),
));
};
match crate::documents::page_text(&path, page) {
Ok(text) => Event::PageTexted { text },
Err(e) => {
usage.duration_ms = started.elapsed().as_millis() as u64;
return Err(fail(error_codes::UNSUPPORTED, e.to_string(), usage.clone()));
}
}
}
Command::PageImage { handle, page } => {
let Some(path) = doc_paths.get(&handle).cloned() else {
usage.duration_ms = started.elapsed().as_millis() as u64;
return Err(fail(
error_codes::CAPABILITY_DENIED,
format!("no such handle: {handle}"),
usage.clone(),
));
};
match crate::documents::render_page(&path, page, RENDER_WIDTH) {
Ok(png) => {
let (handle, len) = handles.insert_bytes(
png,
cuttlefish_abi::MediaKind::Image {
format: "png".into(),
},
);
Event::PageImaged { handle, len }
}
Err(e) => {
usage.duration_ms = started.elapsed().as_millis() as u64;
return Err(fail(error_codes::UNSUPPORTED, e.to_string(), usage.clone()));
}
}
}
Command::Infer {
prompt,
max_tokens,
images,
} => {
if !images.is_empty() && !backend.supports_images() {
usage.duration_ms = started.elapsed().as_millis() as u64;
return Err(fail(
error_codes::UNSUPPORTED,
format!(
"this job supplied {} image(s), but the backend serving `{}` cannot \
accept them. Use a vision-capable model through the `ollama` \
provider, or change the block to send text only.",
images.len(),
backend.model_name()
),
usage.clone(),
));
}
let mut image_bytes = Vec::with_capacity(images.len());
for handle in &images {
match handles.read_all(*handle) {
Ok(bytes) => image_bytes.push(bytes),
Err(e) => {
usage.duration_ms = started.elapsed().as_millis() as u64;
return Err(fail(
error_codes::CAPABILITY_DENIED,
e.to_string(),
usage.clone(),
));
}
}
}
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 request = InferRequest {
prompt: &prompt,
max_tokens,
images: &image_bytes,
};
let infer = backend.infer(request, &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 Err(fail(error_codes::WASM_TRAP, message, usage.clone()));
}
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 Err(cancelled(usage.clone(), "cancelled during inference"));
}
Some(Err(e)) => {
usage.duration_ms = started.elapsed().as_millis() as u64;
return Err(fail(
error_codes::MODEL_LOAD_FAILED,
e.to_string(),
usage.clone(),
));
}
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 Err(fail(
error_codes::WASM_TRAP,
blame(e.to_string()),
usage.clone(),
));
}
};
}
}