use crate::arena::{ContextStack, DataValue, IterGuard};
use crate::node::{PathSegment, ReduceHint};
use crate::opcode::OpCode;
use crate::{CompiledNode, Engine, Result};
use bumpalo::Bump;
use datavalue::NumberValue;
use super::helpers::{
FieldCursor, FusedMapBody, IterArgKind, IterSrc, ResolvedInput, resolve_iter_input,
};
#[inline]
pub(crate) fn evaluate_reduce<'a>(
args: &'a [CompiledNode],
iter_arg_kind: IterArgKind,
ctx: &mut ContextStack<'a>,
engine: &Engine,
arena: &'a Bump,
) -> Result<&'a DataValue<'a>> {
if args.len() < 2 || args.len() > 3 {
return Err(crate::Error::invalid_args());
}
let body = &args[1];
let initial: &'a DataValue<'a> = if args.len() == 3 {
engine.dispatch_node(&args[2], ctx, arena)?
} else {
crate::arena::singletons::singleton_null()
};
if !ctx.is_tracing() && is_map_candidate(&args[0]) {
match try_fused_reduce_map(args, initial, ctx, engine, arena)? {
FusedOutcome::Done(value) => return Ok(value),
FusedOutcome::Bail => {}
}
}
let src = match resolve_iter_input(&args[0], iter_arg_kind, ctx, engine, arena)? {
ResolvedInput::Iterable(s) => s,
ResolvedInput::Empty => return Ok(initial),
ResolvedInput::Bridge(av) => {
return reduce_arena_bridge(av, body, initial, ctx, engine, arena);
}
};
if src.is_empty() {
return Ok(initial);
}
if !ctx.is_tracing() {
if let Some(result) = try_reduce_fast_path(&src, initial, body, arena) {
return Ok(result);
}
}
reduce_general(&src, body, initial, ctx, engine, arena)
}
#[inline]
fn reduce_general<'a>(
src: &IterSrc<'a>,
body: &'a CompiledNode,
initial: &'a DataValue<'a>,
ctx: &mut ContextStack<'a>,
engine: &Engine,
arena: &'a Bump,
) -> Result<&'a DataValue<'a>> {
let len = src.len();
let total = len as u32;
let mut acc_av: &'a DataValue<'a> = initial;
let mut guard = IterGuard::new(ctx);
for i in 0..len {
let item = src.get(i);
guard.step_reduce(item, acc_av);
acc_av = engine.run_iter_body(body, guard.stack(), arena, i as u32, total)?;
}
drop(guard);
Ok(acc_av)
}
#[inline]
fn reduce_arena_bridge<'a>(
input: &'a DataValue<'a>,
body: &'a CompiledNode,
initial: &'a DataValue<'a>,
ctx: &mut ContextStack<'a>,
engine: &Engine,
arena: &'a Bump,
) -> Result<&'a DataValue<'a>> {
match input {
DataValue::Object(pairs) => {
let total = pairs.len() as u32;
let mut acc_av: &'a DataValue<'a> = initial;
let mut guard = IterGuard::new(ctx);
for (i, (_k, v)) in pairs.iter().enumerate() {
guard.step_reduce(v, acc_av);
acc_av = engine.run_iter_body(body, guard.stack(), arena, i as u32, total)?;
}
drop(guard);
Ok(acc_av)
}
_ => Ok(initial),
}
}
enum FusedOutcome<'a> {
Done(&'a DataValue<'a>),
Bail,
}
#[inline(always)]
fn is_map_candidate(node: &CompiledNode) -> bool {
let node = match node {
CompiledNode::Cse(data) => &data.inner,
node => node,
};
matches!(
node,
CompiledNode::BuiltinOperator {
opcode: OpCode::Map,
..
}
)
}
#[inline(never)]
fn try_fused_reduce_map<'a>(
args: &'a [CompiledNode],
initial: &'a DataValue<'a>,
ctx: &mut ContextStack<'a>,
engine: &Engine,
arena: &'a Bump,
) -> Result<FusedOutcome<'a>> {
let map_node = match &args[0] {
CompiledNode::Cse(data) => &data.inner,
node => node,
};
let CompiledNode::BuiltinOperator {
opcode: OpCode::Map,
args: map_args,
iter_arg_kind: map_iter_kind,
..
} = map_node
else {
return Ok(FusedOutcome::Bail);
};
if map_args.len() != 2 {
return Ok(FusedOutcome::Bail);
}
let Some(fold) = detect_fold_shape(&args[1]) else {
return Ok(FusedOutcome::Bail);
};
if !fold.current_segments.is_empty() {
return Ok(FusedOutcome::Bail);
}
let Some(map_body) = FusedMapBody::detect(&map_args[1]) else {
return Ok(FusedOutcome::Bail);
};
let src = match resolve_iter_input(&map_args[0], *map_iter_kind, ctx, engine, arena)? {
ResolvedInput::Iterable(s) => s,
ResolvedInput::Empty => return Ok(FusedOutcome::Done(initial)),
ResolvedInput::Bridge(_) => return Ok(FusedOutcome::Bail),
};
if src.is_empty() {
return Ok(FusedOutcome::Done(initial));
}
Ok(run_fused_fold(&src, initial, &fold, &map_body, arena))
}
fn run_fused_fold<'a>(
src: &IterSrc<'a>,
initial: &'a DataValue<'a>,
fold: &FoldShape<'_>,
map_body: &FusedMapBody<'_>,
arena: &'a Bump,
) -> FusedOutcome<'a> {
let op = fold.op;
let acc_is_lhs = fold.acc_is_lhs;
let lit = match map_body {
FusedMapBody::ArithVarLit { lit, .. } => match lit.as_number() {
Some(n) => *n,
None => return FusedOutcome::Bail,
},
_ => NumberValue::from_i64(0),
};
let mut cursors = FusedCursors::new(map_body);
let Some(mut acc) = initial.as_number().copied() else {
return FusedOutcome::Bail;
};
for i in 0..src.len() {
let item = src.get(i);
let Some(mapped) = mapped_number(map_body, &mut cursors, item, lit) else {
return FusedOutcome::Bail;
};
let Some(next) = fold_number(op, acc_is_lhs, acc, mapped) else {
return FusedOutcome::Bail;
};
acc = next;
}
FusedOutcome::Done(alloc_number(arena, acc))
}
#[inline(always)]
fn alloc_number<'a>(arena: &'a Bump, n: NumberValue) -> &'a DataValue<'a> {
if let NumberValue::Integer(i) = n {
if let Some(singleton) = crate::arena::singletons::singleton_small_int(i) {
return singleton;
}
}
arena.alloc(DataValue::Number(n))
}
struct FusedCursors<'n> {
a: FieldCursor<'n>,
b: Option<FieldCursor<'n>>,
}
impl<'n> FusedCursors<'n> {
fn new(map_body: &FusedMapBody<'n>) -> Self {
match map_body {
FusedMapBody::Extract { segments } => Self {
a: FieldCursor::new(segments),
b: None,
},
FusedMapBody::ArithVarLit { segments, .. } => Self {
a: FieldCursor::new(segments),
b: None,
},
FusedMapBody::ArithVarVar {
a_segments,
b_segments,
..
} => Self {
a: FieldCursor::new(a_segments),
b: Some(FieldCursor::new(b_segments)),
},
}
}
}
#[inline(always)]
fn mapped_number<'a>(
map_body: &FusedMapBody<'_>,
cursors: &mut FusedCursors<'_>,
item: &'a DataValue<'a>,
lit: NumberValue,
) -> Option<NumberValue> {
match map_body {
FusedMapBody::Extract { .. } => cursors.a.resolve(item)?.as_number().copied(),
FusedMapBody::ArithVarLit { op, var_is_lhs, .. } => {
let v = *cursors.a.resolve(item)?.as_number()?;
let (x, y) = if *var_is_lhs { (v, lit) } else { (lit, v) };
arith_number(*op, x, y)
}
FusedMapBody::ArithVarVar { op, .. } => {
let a = *cursors.a.resolve(item)?.as_number()?;
let b = *cursors.b.as_mut()?.resolve(item)?.as_number()?;
arith_number(*op, a, b)
}
}
}
type IntOp = fn(i64, i64) -> Option<i64>;
type FloatOp = fn(f64, f64) -> f64;
#[inline(always)]
fn arith_number(op: OpCode, a: NumberValue, b: NumberValue) -> Option<NumberValue> {
let (int_op, float_op): (IntOp, FloatOp) = match op {
OpCode::Add => (i64::checked_add, |x, y| x + y),
OpCode::Subtract => (i64::checked_sub, |x, y| x - y),
OpCode::Multiply => (i64::checked_mul, |x, y| x * y),
_ => return None,
};
Some(match (a.as_i64(), b.as_i64()) {
(Some(x), Some(y)) => crate::operators::arithmetic::try_int_op(x, y, int_op, float_op),
_ => NumberValue::from_f64(float_op(a.as_f64(), b.as_f64())),
})
}
#[inline(always)]
fn fold_number(
op: OpCode,
acc_is_lhs: bool,
acc: NumberValue,
cur: NumberValue,
) -> Option<NumberValue> {
if acc_is_lhs {
arith_number(op, acc, cur)
} else {
arith_number(op, cur, acc)
}
}
struct FoldShape<'a> {
op: OpCode,
acc_is_lhs: bool,
current_segments: &'a [PathSegment],
}
fn detect_fold_shape(body: &CompiledNode) -> Option<FoldShape<'_>> {
let (opcode, body_args) = match body {
CompiledNode::BuiltinOperator { opcode, args, .. } => (*opcode, args),
_ => return None,
};
if body_args.len() != 2 || !matches!(opcode, OpCode::Add | OpCode::Multiply | OpCode::Subtract)
{
return None;
}
let (current_arg, acc_is_lhs) = match (&body_args[0], &body_args[1]) {
(
CompiledNode::Var {
reduce_hint: hint0, ..
},
CompiledNode::Var {
reduce_hint: hint1, ..
},
) => match (hint0, hint1) {
(
ReduceHint::Current | ReduceHint::CurrentPath,
ReduceHint::Accumulator | ReduceHint::AccumulatorPath,
) => (&body_args[0], false),
(
ReduceHint::Accumulator | ReduceHint::AccumulatorPath,
ReduceHint::Current | ReduceHint::CurrentPath,
) => (&body_args[1], true),
_ => return None,
},
_ => return None,
};
let current_segments = if let CompiledNode::Var {
segments,
reduce_hint,
..
} = current_arg
{
match reduce_hint {
ReduceHint::Current => &[][..],
ReduceHint::CurrentPath if segments.len() >= 2 => &segments[1..],
_ => return None,
}
} else {
return None;
};
Some(FoldShape {
op: opcode,
acc_is_lhs,
current_segments,
})
}
fn try_reduce_fast_path<'a>(
src: &IterSrc<'a>,
initial: &'a DataValue<'a>,
body: &CompiledNode,
arena: &'a Bump,
) -> Option<&'a DataValue<'a>> {
let FoldShape {
op,
acc_is_lhs,
current_segments,
} = detect_fold_shape(body)?;
let mut current_field = FieldCursor::new(current_segments);
let mut acc = initial.as_number().copied()?;
for i in 0..src.len() {
let item = src.get(i);
let cur = *current_field.resolve(item)?.as_number()?;
acc = fold_number(op, acc_is_lhs, acc, cur)?;
}
Some(alloc_number(arena, acc))
}