use std::path::PathBuf;
use std::sync::Arc;
use fv_compute::{RegistryError, Runtime};
use fv_compute_container::ContainerBackend;
use fv_compute_wasm::WasmBackend;
pub use fv_plan::KineticsRunner as WorkerCompute;
pub fn with_roots(roots: &[PathBuf]) -> Result<WorkerCompute, RegistryError> {
Ok(WorkerCompute::new(runtime_with_roots(roots)?))
}
pub fn runtime_with_roots(roots: &[PathBuf]) -> Result<Arc<Runtime>, RegistryError> {
let mut builder = Runtime::builder();
for root in roots {
builder = builder.root(root.clone());
}
let wasm = WasmBackend::new().map_err(|e| RegistryError::Io {
path: PathBuf::from("<wasm-engine>"),
msg: e.to_string(),
})?;
Ok(Arc::new(
builder.backend(wasm).backend(ContainerBackend::new()).build()?,
))
}
pub fn run_batch(
runtime: &Runtime,
step: &serde_json::Value,
batch: &arrow::array::RecordBatch,
) -> Result<arrow::array::RecordBatch, String> {
let selector = WorkerCompute::selector(step)?;
let compute = runtime.load(selector).map_err(|e| e.to_string())?;
let manifest = compute.manifest();
if manifest.inputs.len() > 1 {
return Err(format!(
"transform `{selector}` declares {} inputs; a stream stage supplies one",
manifest.inputs.len()
));
}
let typed = match manifest.inputs.first() {
Some(signature) => typed_batch(&signature.to_arrow_schema_ref(), batch)?,
None => batch.clone(),
};
compute.run(&[typed]).map_err(|e| e.to_string())
}
pub fn typed_batch(
schema: &arrow::datatypes::SchemaRef,
batch: &arrow::array::RecordBatch,
) -> Result<arrow::array::RecordBatch, String> {
let n = batch.num_rows();
let opts = arrow::compute::CastOptions {
safe: true,
..Default::default()
};
let cols: Vec<arrow::array::ArrayRef> = schema
.fields()
.iter()
.map(|f| match batch.column_by_name(f.name()) {
Some(c) if c.data_type() == f.data_type() => Ok(Arc::clone(c)),
Some(c) => arrow::compute::cast_with_options(c, f.data_type(), &opts).map_err(|e| e.to_string()),
None => Ok(arrow::array::new_null_array(f.data_type(), n)),
})
.collect::<Result<_, _>>()?;
let relaxed = Arc::new(arrow::datatypes::Schema::new(
schema
.fields()
.iter()
.map(|f| f.as_ref().clone().with_nullable(true))
.collect::<Vec<_>>(),
));
arrow::array::RecordBatch::try_new(relaxed, cols).map_err(|e| e.to_string())
}
pub fn from_env() -> Result<WorkerCompute, RegistryError> {
let roots: Vec<PathBuf> = std::env::var("FV_TRANSFORMS_DIR")
.unwrap_or_default()
.split(':')
.filter(|s| !s.is_empty())
.map(PathBuf::from)
.collect();
with_roots(&roots)
}
#[cfg(test)]
mod tests {
use super::*;
use fv_plan::build::ComputeRunner;
use fv_plan::row::Row;
use fv_value::Value;
fn wasm_fixtures() -> PathBuf {
PathBuf::from(env!("CARGO_MANIFEST_DIR")).join("tests/fixtures")
}
#[tokio::test]
async fn runs_a_wasm_transform_via_the_registry() {
let runner = with_roots(&[wasm_fixtures()]).unwrap();
let step = serde_json::json!({ "op": "wasm", "ref": "spikeTotal" });
let inputs = vec![(
"raw".to_string(),
vec![Row(vec![
("id".into(), Value::Num(100.0)),
("amount".into(), Value::Num(7.25)),
])],
)];
let out = runner.run(&step, &inputs).await.unwrap();
assert_eq!(out[0].get("total"), Value::Num(114.5)); }
#[tokio::test]
async fn runs_a_container_transform_via_the_registry() {
let runner = with_roots(&[wasm_fixtures()]).unwrap();
let step = serde_json::json!({ "op": "container", "ref": "riskBand" });
let inputs = vec![(
"staged".to_string(),
vec![
Row(vec![
("customerId".into(), Value::Str("C1".into())),
("creditTerms".into(), Value::Str("NET90".into())),
("creditLimit".into(), Value::Num(90000.0)),
]),
Row(vec![
("customerId".into(), Value::Str("C2".into())),
("creditTerms".into(), Value::Str("PREPAID".into())),
("creditLimit".into(), Value::Num(1000.0)),
]),
],
)];
let out = runner.run(&step, &inputs).await.unwrap();
assert_eq!(out.len(), 2);
assert_eq!(out[0].get("customerId"), Value::Str("C1".into()));
assert_eq!(out[0].get("riskBand"), Value::Str("HIGH".into())); assert_eq!(out[1].get("riskBand"), Value::Str("LOW".into())); }
#[test]
fn run_batch_types_the_batch_to_the_signature_and_runs_the_transform() {
use arrow::array::{Float64Array, Int64Array, RecordBatch, StringArray};
use arrow::datatypes::{DataType, Field, Schema};
let rt = runtime_with_roots(&[wasm_fixtures()]).unwrap();
let step = serde_json::json!({ "op": "wasm", "ref": "spikeTotal" });
let schema = Arc::new(Schema::new(vec![
Field::new("id", DataType::Float64, true),
Field::new("amount", DataType::Float64, true),
Field::new("extra", DataType::Utf8, true),
]));
let b = RecordBatch::try_new(
schema,
vec![
Arc::new(Float64Array::from(vec![100.0, 200.0])),
Arc::new(Float64Array::from(vec![7.25, 1.0])),
Arc::new(StringArray::from(vec!["x", "y"])),
],
)
.unwrap();
let out = run_batch(&rt, &step, &b).unwrap();
let total = out.column_by_name("total").unwrap();
let totals: Vec<f64> = match total.data_type() {
DataType::Float64 => total.as_any().downcast_ref::<Float64Array>().unwrap().values().to_vec(),
_ => total
.as_any()
.downcast_ref::<Int64Array>()
.map(|a| a.values().iter().map(|v| *v as f64).collect())
.unwrap_or_default(),
};
assert_eq!(totals, vec![114.5, 202.0]); }
#[tokio::test]
async fn unknown_transform_fails_closed() {
let runner = with_roots(&[wasm_fixtures()]).unwrap();
let step = serde_json::json!({ "op": "wasm", "ref": "doesNotExist" });
let err = runner.run(&step, &[("raw".into(), vec![])]).await.unwrap_err();
assert!(err.contains("no such transform"), "got: {err}");
}
}