use crate::{
args::{ArgValues, FromArgs},
bytecode::VM,
exception_private::{ExcType, ExcTypeExt, RunResult},
heap::{DropGuard, DropWithContext, HeapData, HeapId},
intern::StaticStrings,
modules::ModuleFunctions,
types::{
ItertoolsIter, Module, Type,
itertools::{Chain, Compress, Count, Cycle, Islice, Pairwise, Repeat},
},
value::Value,
};
#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash, strum::Display, serde::Serialize, serde::Deserialize)]
#[strum(serialize_all = "snake_case")]
pub(crate) enum ItertoolsFunctions {
Count,
Repeat,
Pairwise,
Compress,
Islice,
Chain,
Cycle,
}
const ITERTOOLS_FUNCTIONS: &[(StaticStrings, ItertoolsFunctions)] = &[
(StaticStrings::Count, ItertoolsFunctions::Count),
(StaticStrings::Repeat, ItertoolsFunctions::Repeat),
(StaticStrings::Pairwise, ItertoolsFunctions::Pairwise),
(StaticStrings::Compress, ItertoolsFunctions::Compress),
(StaticStrings::Islice, ItertoolsFunctions::Islice),
(StaticStrings::Chain, ItertoolsFunctions::Chain),
(StaticStrings::Cycle, ItertoolsFunctions::Cycle),
];
pub fn create_module(vm: &mut VM<'_>) -> HeapId {
let mut module = Module::new(StaticStrings::Itertools);
for (name, func) in ITERTOOLS_FUNCTIONS {
module.set_attr(*name, Value::ModuleFunction(ModuleFunctions::Itertools(*func)), vm);
}
vm.heap.allocate(HeapData::Module(Box::new(module)))
}
pub(super) fn call(vm: &mut VM<'_>, function: ItertoolsFunctions, args: ArgValues) -> RunResult<Value> {
match function {
ItertoolsFunctions::Count => call_count(vm, args),
ItertoolsFunctions::Repeat => call_repeat(vm, args),
ItertoolsFunctions::Pairwise => call_pairwise(vm, args),
ItertoolsFunctions::Compress => call_compress(vm, args),
ItertoolsFunctions::Islice => call_islice(vm, args),
ItertoolsFunctions::Chain => call_chain(vm, args),
ItertoolsFunctions::Cycle => call_cycle(vm, args),
}
}
#[derive(FromArgs)]
#[from_args(name = "count", style = c_named, at_most_total)]
struct CountArgs {
#[from_args(default = Value::Int(0))]
start: Value,
#[from_args(default = Value::Int(1))]
step: Value,
}
fn call_count(vm: &mut VM<'_>, args: ArgValues) -> RunResult<Value> {
let CountArgs { start, step } = CountArgs::from_args(args, vm)?;
if is_number(&start, vm) && is_number(&step, vm) {
let iter = ItertoolsIter::Count(Count::new(normalize_bool(start), normalize_bool(step)));
Ok(Value::Ref(vm.heap.allocate(HeapData::Itertools(iter))))
} else {
start.drop_with(vm);
step.drop_with(vm);
Err(ExcType::type_error("a number is required"))
}
}
#[derive(FromArgs)]
#[from_args(name = "repeat", style = c_named, at_most_total)]
struct RepeatArgs {
object: Value,
#[from_args(default)]
times: Option<Value>,
}
fn call_repeat(vm: &mut VM<'_>, args: ArgValues) -> RunResult<Value> {
let RepeatArgs { object, times } = RepeatArgs::from_args(args, vm)?;
let remaining = match times {
None => None,
Some(times) => {
let count = repeat_times(×, vm);
times.drop_with(vm);
match count {
Ok(count) => Some(count),
Err(error) => {
object.drop_with(vm);
return Err(error);
}
}
}
};
let iter = ItertoolsIter::Repeat(Repeat::new(object, remaining));
Ok(Value::Ref(vm.heap.allocate(HeapData::Itertools(iter))))
}
fn is_number(value: &Value, vm: &VM<'_>) -> bool {
matches!(value.py_type_heap(vm.heap), Type::Int | Type::Float | Type::Bool)
}
fn normalize_bool(value: Value) -> Value {
match value {
Value::Bool(b) => Value::Int(i64::from(b)),
other => other,
}
}
fn repeat_times(value: &Value, vm: &VM<'_>) -> RunResult<usize> {
let count = match value {
Value::Bool(b) => i64::from(*b),
other => other.as_int(vm)?,
};
Ok(usize::try_from(count.max(0)).unwrap_or(usize::MAX))
}
#[derive(FromArgs)]
#[from_args(name = "pairwise", style = unpack)]
struct PairwiseArgs {
#[from_args(pos_only)]
iterable: Value,
}
fn call_pairwise(vm: &mut VM<'_>, args: ArgValues) -> RunResult<Value> {
let PairwiseArgs { iterable } = PairwiseArgs::from_args(args, vm)?;
let source = iterable.into_py_iter(vm)?;
let iter = ItertoolsIter::Pairwise(Pairwise::new(source));
Ok(Value::Ref(vm.heap.allocate(HeapData::Itertools(iter))))
}
#[derive(FromArgs)]
#[from_args(name = "compress", style = c_named, at_most_total)]
struct CompressArgs {
data: Value,
selectors: Value,
}
fn call_compress(vm: &mut VM<'_>, args: ArgValues) -> RunResult<Value> {
let CompressArgs { data, selectors } = CompressArgs::from_args(args, vm)?;
let mut guard = DropGuard::new(selectors, vm);
let data = data.into_py_iter(guard.ctx())?;
let (selectors, vm) = guard.into_parts();
let mut guard = DropGuard::new(data, vm);
let selectors = selectors.into_py_iter(guard.ctx())?;
let (data, vm) = guard.into_parts();
let iter = ItertoolsIter::Compress(Compress::new(data, selectors));
Ok(Value::Ref(vm.heap.allocate(HeapData::Itertools(iter))))
}
#[derive(FromArgs)]
#[from_args(name = "islice", style = unpack)]
struct IsliceArgs {
#[from_args(pos_only)]
iterable: Value,
#[from_args(pos_only)]
first: Value,
#[from_args(pos_only, default)]
second: Option<Value>,
#[from_args(pos_only, default)]
third: Option<Value>,
}
fn call_islice(vm: &mut VM<'_>, args: ArgValues) -> RunResult<Value> {
let IsliceArgs {
iterable,
first,
second,
third,
} = IsliceArgs::from_args(args, vm)?;
let mut guard = DropGuard::new(iterable, vm);
let vm = guard.ctx();
let bounds = islice_bounds(&first, second.as_ref(), third.as_ref(), vm);
first.drop_with(vm);
second.drop_with(vm);
third.drop_with(vm);
let (start, stop, step) = bounds?;
let (iterable, vm) = guard.into_parts();
let source = iterable.into_py_iter(vm)?;
let iter = ItertoolsIter::Islice(Islice::new(source, start, stop, step));
Ok(Value::Ref(vm.heap.allocate(HeapData::Itertools(iter))))
}
fn islice_bounds(
first: &Value,
second: Option<&Value>,
third: Option<&Value>,
vm: &VM<'_>,
) -> RunResult<(usize, Option<usize>, usize)> {
match second {
None => match islice_index(first, vm) {
IsliceBound::Unbounded => Ok((0, None, 1)),
IsliceBound::Index(stop) => Ok((0, Some(stop), 1)),
IsliceBound::Invalid => Err(ExcType::islice_bad_stop()),
},
Some(second) => {
let start = match islice_index(first, vm) {
IsliceBound::Unbounded => 0,
IsliceBound::Index(start) => start,
IsliceBound::Invalid => return Err(ExcType::islice_bad_indices()),
};
let stop = match islice_index(second, vm) {
IsliceBound::Unbounded => None,
IsliceBound::Index(stop) => Some(stop),
IsliceBound::Invalid => return Err(ExcType::islice_bad_indices()),
};
let step = match third.map(|third| islice_index(third, vm)) {
None | Some(IsliceBound::Unbounded) => 1,
Some(IsliceBound::Index(step)) if step > 0 => step,
Some(_) => return Err(ExcType::islice_bad_step()),
};
Ok((start, stop, step))
}
}
}
enum IsliceBound {
Unbounded,
Index(usize),
Invalid,
}
fn islice_index(value: &Value, vm: &VM<'_>) -> IsliceBound {
let index = match value {
Value::None => return IsliceBound::Unbounded,
Value::Bool(b) => i64::from(*b),
other => match other.as_int(vm) {
Ok(index) => index,
Err(_) => return IsliceBound::Invalid,
},
};
usize::try_from(index).map_or(IsliceBound::Invalid, IsliceBound::Index)
}
fn call_chain(vm: &mut VM<'_>, args: ArgValues) -> RunResult<Value> {
let iterables: Vec<Value> = args.into_pos_only("chain", vm.heap)?.collect();
let iter = ItertoolsIter::Chain(Chain::new(iterables));
Ok(Value::Ref(vm.heap.allocate(HeapData::Itertools(iter))))
}
#[derive(FromArgs)]
#[from_args(name = "cycle", style = unpack)]
struct CycleArgs {
#[from_args(pos_only)]
iterable: Value,
}
fn call_cycle(vm: &mut VM<'_>, args: ArgValues) -> RunResult<Value> {
let CycleArgs { iterable } = CycleArgs::from_args(args, vm)?;
let source = iterable.into_py_iter(vm)?;
let iter = ItertoolsIter::Cycle(Cycle::new(source));
Ok(Value::Ref(vm.heap.allocate(HeapData::Itertools(iter))))
}