#![cfg(feature = "budget")]
use async_trait::async_trait;
use dataflow_rs::engine::functions::AsyncFunctionHandler;
use dataflow_rs::engine::message::Message;
use dataflow_rs::{
DataflowError, Engine, Result, TaskContext, TaskOutcome, Template, TemplateCompiler, Workflow,
};
use serde_json::{Value, json};
mod common;
use common::workflow;
const N: usize = 20_000;
const TIGHT: u64 = 2_000;
struct SumViaEval;
#[async_trait]
impl AsyncFunctionHandler for SumViaEval {
type Input = SumInput;
fn compile_input(input: &mut Self::Input, c: &TemplateCompiler) -> Result<()> {
input.expr.compile(c, "expr")
}
async fn execute(&self, ctx: &mut TaskContext<'_>, input: &Self::Input) -> Result<TaskOutcome> {
let total = input.expr.eval(ctx)?;
ctx.set("data.total", total);
Ok(TaskOutcome::Success)
}
}
#[derive(serde::Deserialize)]
struct SumInput {
expr: Template,
}
fn fold_expr() -> Value {
json!({
"reduce": [
{ "var": "data.xs" },
{ "+": [{ "var": "current" }, { "var": "accumulator" }] },
0
]
})
}
fn sum_workflow(expr: Value) -> Workflow {
workflow(json!({
"id": "w", "name": "w", "priority": 0,
"tasks": [
{ "id": "sum", "name": "sum", "function": {
"name": "sum_via_eval",
"input": { "expr": expr } } }
]
}))
}
fn engine(expr: Value, budget: Option<u64>) -> Engine {
let builder = Engine::builder()
.register("sum_via_eval", SumViaEval)
.with_workflow(sum_workflow(expr));
match budget {
Some(b) => builder.with_ops_budget(b),
None => builder,
}
.build()
.unwrap()
}
fn message() -> Message {
Message::builder()
.data_json(&json!({ "xs": vec![1; N] }))
.build()
}
fn total(m: &Message) -> Value {
Value::from(m.data().get("total").unwrap())
}
fn task_error(m: &Message) -> &dataflow_rs::ErrorInfo {
m.errors()
.iter()
.find(|e| e.task_id.is_some())
.unwrap_or_else(|| panic!("no task-level error recorded, got {:?}", m.errors()))
}
#[tokio::test]
async fn an_expression_over_the_ceiling_is_refused_with_its_own_error_code() {
let engine = engine(fold_expr(), Some(TIGHT));
let mut m = message();
let outcome = engine.process_message(&mut m).await;
assert!(
outcome.is_err(),
"a task whose only evaluation was refused should not report success"
);
let err = task_error(&m);
assert_eq!(
err.code, "BUDGET_EXCEEDED",
"a budget abort must not collapse into LOGIC_ERROR: {err:?}"
);
assert!(
m.data().get("total").is_none_or(|v| v.is_null()),
"the write must not happen — the ceiling aborts *before* the work"
);
assert_eq!(
err.message.matches("Operation budget exceeded").count(),
1,
"datalogic's message already names the condition; the variant must \
not prefix it again: {}",
err.message
);
}
#[tokio::test]
async fn the_same_expression_is_fine_unbudgeted() {
let engine = engine(fold_expr(), None);
let mut m = message();
engine.process_message(&mut m).await.unwrap();
assert!(m.errors().is_empty(), "unbudgeted run: {:?}", m.errors());
assert_eq!(total(&m), json!(N), "the fold should have completed");
}
#[tokio::test]
async fn a_cheap_expression_passes_under_the_same_ceiling() {
let engine = engine(json!({ "+": [1, 2] }), Some(TIGHT));
let mut m = Message::builder().build();
engine.process_message(&mut m).await.unwrap();
assert!(m.errors().is_empty(), "{:?}", m.errors());
assert_eq!(total(&m), json!(3));
}
#[tokio::test]
async fn the_ceiling_survives_a_hot_reload() {
let engine = engine(json!({ "+": [1, 2] }), Some(TIGHT));
let reloaded = engine
.with_new_workflows(vec![sum_workflow(fold_expr())])
.unwrap();
let mut m = message();
let outcome = reloaded.process_message(&mut m).await;
assert!(outcome.is_err(), "the reloaded engine dropped its ceiling");
assert_eq!(task_error(&m).code, "BUDGET_EXCEEDED");
}
#[test]
fn a_budget_error_is_not_retryable() {
assert!(!DataflowError::BudgetExceeded("over".into()).retryable());
}
#[test]
fn a_budget_error_round_trips_through_serde() {
let err = DataflowError::BudgetExceeded("over at node 3".into());
let json = serde_json::to_string(&err).unwrap();
let back: DataflowError = serde_json::from_str(&json).unwrap();
assert!(matches!(back, DataflowError::BudgetExceeded(m) if m == "over at node 3"));
}