use std::{cmp::Ordering, mem};
use smallvec::smallvec;
use crate::{
args::{ArgValues, KwargsValues},
bytecode::VM,
defer_drop, defer_drop_mut,
exception_private::{ExcType, ExcTypeExt, RunResult},
heap::{DropGuard, DropWithContext, HeapData, HeapId, HeapReadOutput},
resource_checks::check_repeat_size,
types::{Dict, List, PyTrait, allocate_tuple, iter::collect_owned_iterable, py_trait::CmpOrder},
value::{VALUE_SIZE, Value},
};
fn count_arith(lhs: &Value, rhs: &Value, subtract: bool, vm: &mut VM<'_>) -> RunResult<Value> {
if subtract {
lhs.py_sub(rhs, vm)
} else {
lhs.py_add(rhs, vm)
}
}
#[derive(Clone, Copy)]
enum CountCmp {
Le,
Lt,
Ge,
Gt,
}
impl CountCmp {
fn symbol(self) -> &'static str {
match self {
Self::Le => "<=",
Self::Lt => "<",
Self::Ge => ">=",
Self::Gt => ">",
}
}
fn holds(self, ordering: Ordering) -> bool {
match self {
Self::Le => ordering != Ordering::Greater,
Self::Lt => ordering == Ordering::Less,
Self::Ge => ordering != Ordering::Less,
Self::Gt => ordering == Ordering::Greater,
}
}
}
fn count_holds(lhs: &Value, rhs: &Value, op: CountCmp, vm: &mut VM<'_>) -> RunResult<bool> {
match lhs.py_cmp(rhs, vm)? {
CmpOrder::Ordered(ordering) => Ok(op.holds(ordering)),
CmpOrder::Unordered => Ok(false),
CmpOrder::Incomparable => Err(ExcType::type_error_ordering(
op.symbol(),
&lhs.py_type_name(vm),
&rhs.py_type_name(vm),
)),
}
}
fn count_sort_cmp(lhs: &Value, rhs: &Value, vm: &mut VM<'_>) -> RunResult<Ordering> {
match lhs.py_cmp(rhs, vm)? {
CmpOrder::Ordered(ordering) => Ok(ordering),
CmpOrder::Unordered => Ok(Ordering::Equal),
CmpOrder::Incomparable => Err(ExcType::type_error_ordering(
"<",
&lhs.py_type_name(vm),
&rhs.py_type_name(vm),
)),
}
}
fn count_is_positive(value: &Value, vm: &mut VM<'_>) -> RunResult<bool> {
count_holds(value, &Value::Int(0), CountCmp::Gt, vm)
}
fn count_repeat_len(value: &Value, vm: &mut VM<'_>) -> RunResult<usize> {
match value {
Value::Int(i) => Ok(usize::try_from(*i).unwrap_or(0)),
Value::Bool(b) => Ok(usize::from(*b)),
Value::Ref(id) if matches!(vm.heap.get(*id), HeapData::LongInt(_)) => Err(ExcType::overflow_c_ssize_t()),
Value::InternLongInt(_) => Err(ExcType::overflow_c_ssize_t()),
other => Err(ExcType::type_error_not_an_integer(&other.py_type_name(vm))),
}
}
fn counter_bump(
dict_id: HeapId,
key: Value,
delta: &Value,
subtract: bool,
delta_first: bool,
vm: &mut VM<'_>,
) -> RunResult<()> {
let HeapReadOutput::Dict(mut dict) = vm.heap.read(dict_id) else {
unreachable!("counter_bump on a non-dict heap entry");
};
let mut key_guard = DropGuard::new(key, vm);
let (key, vm) = key_guard.as_parts_mut();
let base = dict.dict_get(key, vm)?.unwrap_or(Value::Int(0));
let total = if subtract || !delta_first {
count_arith(&base, delta, subtract, vm)
} else {
count_arith(delta, &base, false, vm)
};
base.drop_with(vm);
let total = total?;
let (key, vm) = key_guard.into_parts();
if let Some(old) = dict.set(key, total, vm)? {
old.drop_with(vm);
}
Ok(())
}
pub(crate) fn counter_update(
dict_id: HeapId,
source: Option<Value>,
kwargs: KwargsValues,
subtract: bool,
vm: &mut VM<'_>,
) -> RunResult<()> {
if let Some(source) = source
&& !matches!(source, Value::None)
{
let is_mapping = matches!(&source, Value::Ref(id) if matches!(vm.heap.get(*id), HeapData::Dict(_)));
if is_mapping {
let Value::Ref(src_id) = source else {
unreachable!("mapping is a ref");
};
let pairs: Vec<(Value, Value)> = {
let HeapReadOutput::Dict(src) = vm.heap.read(src_id) else {
unreachable!("mapping is a dict");
};
let src = src.get(vm.heap);
src.iter()
.map(|(k, v)| (k.clone_with_heap(vm.heap), v.clone_with_heap(vm.heap)))
.collect()
};
source.drop_with(vm);
counter_merge_mapping(dict_id, pairs, subtract, vm)?;
} else {
let items = collect_owned_iterable::<Vec<Value>>(source, vm)?.into_iter();
defer_drop_mut!(items, vm);
for item in items.by_ref() {
counter_bump(dict_id, item, &Value::Int(1), subtract, false, vm)?;
}
}
}
let kwarg_pairs = counter_kwarg_deltas(kwargs);
counter_merge_mapping(dict_id, kwarg_pairs, subtract, vm)?;
Ok(())
}
fn counter_merge_mapping(
dict_id: HeapId,
pairs: Vec<(Value, Value)>,
subtract: bool,
vm: &mut VM<'_>,
) -> RunResult<()> {
let is_empty = {
let HeapReadOutput::Dict(dict) = vm.heap.read(dict_id) else {
unreachable!("counter_merge_mapping on a non-dict heap entry");
};
dict.get(vm.heap).is_empty()
};
let pairs = pairs.into_iter();
defer_drop_mut!(pairs, vm);
for (key, count) in pairs.by_ref() {
if is_empty && !subtract {
counter_set(dict_id, key, count, vm)?;
} else {
let outcome = counter_bump(dict_id, key, &count, subtract, true, vm);
count.drop_with(vm);
outcome?;
}
}
Ok(())
}
fn counter_kwarg_deltas(kwargs: KwargsValues) -> Vec<(Value, Value)> {
match kwargs {
KwargsValues::Empty => Vec::new(),
KwargsValues::Inline(kvs) => kvs.into_iter().map(|(id, v)| (Value::InternString(id), v)).collect(),
KwargsValues::Pairs(kvs) => kvs,
KwargsValues::Dict(dict) => dict.into_iter().collect(),
}
}
pub(crate) fn counter_update_method(
dict_id: HeapId,
args: ArgValues,
subtract: bool,
vm: &mut VM<'_>,
) -> RunResult<Value> {
let (mut pos, kwargs) = args.into_parts();
let total = pos.len();
let source = (total > 0).then(|| pos.next().expect("len checked"));
let name = if subtract { "Counter.subtract" } else { "Counter.update" };
if total > 1 {
source.drop_with(vm);
pos.drop_with(vm);
kwargs.drop_with(vm);
return Err(ExcType::type_error_too_many_positional_range(name, 1, 2, total + 1, 0));
}
counter_update(dict_id, source, kwargs, subtract, vm)?;
Ok(Value::None)
}
pub(crate) fn counter_total(dict_id: HeapId, vm: &mut VM<'_>) -> RunResult<Value> {
let counts = counter_count_snapshot(dict_id, vm).into_iter();
defer_drop_mut!(counts, vm);
let mut sum_guard = DropGuard::new(Value::Int(0), vm);
for count in counts.by_ref() {
let (sum, vm) = sum_guard.as_parts_mut();
let next = count_arith(sum, &count, false, vm);
count.drop_with(vm);
mem::replace(sum, next?).drop_with(vm);
}
Ok(sum_guard.into_inner())
}
pub(crate) fn counter_order(counts: Vec<Value>, vm: &mut VM<'_>) -> RunResult<Vec<usize>> {
let mut order: Vec<usize> = (0..counts.len()).collect();
let mut failure = None;
order.sort_by(|&a, &b| {
if failure.is_some() {
return Ordering::Equal;
}
match count_sort_cmp(&counts[b], &counts[a], vm) {
Ok(ordering) => ordering,
Err(e) => {
failure = Some(e);
Ordering::Equal
}
}
});
counts.drop_with(vm);
match failure {
Some(e) => Err(e),
None => Ok(order),
}
}
fn counter_count_snapshot(dict_id: HeapId, vm: &mut VM<'_>) -> Vec<Value> {
let HeapReadOutput::Dict(dict) = vm.heap.read(dict_id) else {
unreachable!("counter_count_snapshot on a non-dict heap entry");
};
let dict = dict.get(vm.heap);
dict.iter().map(|(_, v)| v.clone_with_heap(vm.heap)).collect()
}
pub(crate) fn counter_most_common(dict_id: HeapId, args: ArgValues, vm: &mut VM<'_>) -> RunResult<Value> {
let n_arg = args.get_zero_one_arg("most_common", vm.heap)?;
let limit = match n_arg {
None | Some(Value::None) => None,
Some(Value::Int(n)) => Some(usize::try_from(n.max(0)).unwrap_or(usize::MAX)),
Some(Value::Bool(b)) => Some(usize::from(b)),
Some(value @ Value::Ref(id)) if matches!(vm.heap.get(id), HeapData::LongInt(_)) => {
let negative = match vm.heap.get(id) {
HeapData::LongInt(li) => li.is_negative(),
_ => unreachable!("guarded by the matches! above"),
};
value.drop_with(vm);
Some(if negative { 0 } else { usize::MAX })
}
Some(other) => {
let ty = other.py_type_name(vm);
other.drop_with(vm);
return Err(ExcType::type_error(format!(
"'{ty}' object cannot be interpreted as an integer"
)));
}
};
let order = counter_order(counter_count_snapshot(dict_id, vm), vm)?;
let take = limit.unwrap_or(order.len()).min(order.len());
let mut items: Vec<Value> = Vec::with_capacity(take);
for &i in &order[..take] {
let HeapReadOutput::Dict(dict) = vm.heap.read(dict_id) else {
unreachable!("counter_most_common on a non-dict heap entry");
};
let dict = dict.get(vm.heap);
let key = dict.key_at(i).expect("index in range").clone_with_heap(vm.heap);
let count = dict.value_at(i).expect("index in range").clone_with_heap(vm.heap);
let pair = allocate_tuple(smallvec![key, count], vm.heap);
items.push(pair);
}
Ok(Value::Ref(vm.heap.allocate(HeapData::List(List::new(items)))))
}
pub(crate) fn counter_elements(dict_id: HeapId, args: ArgValues, vm: &mut VM<'_>) -> RunResult<Value> {
args.check_zero_args("elements", vm.heap)?;
let HeapReadOutput::Dict(dict) = vm.heap.read(dict_id) else {
unreachable!("counter_elements on a non-dict heap entry");
};
let n = dict.get(vm.heap).len();
let mut lengths = Vec::with_capacity(n);
for i in 0..n {
let count = dict
.get(vm.heap)
.value_at(i)
.expect("index in range")
.clone_with_heap(vm.heap);
let len = count_repeat_len(&count, vm);
count.drop_with(vm);
lengths.push(len?);
}
let total = lengths.iter().fold(0usize, |acc, len| acc.saturating_add(*len));
check_repeat_size(VALUE_SIZE, total, vm.heap.tracker())?;
let mut items: Vec<Value> = Vec::new();
for (i, &count) in lengths.iter().enumerate() {
for _ in 0..count {
let key = dict
.get(vm.heap)
.key_at(i)
.expect("index in range")
.clone_with_heap(vm.heap);
items.push(key);
}
}
Ok(Value::Ref(vm.heap.allocate(HeapData::List(List::new(items)))))
}
#[derive(Clone, Copy)]
pub(crate) enum CounterOp {
Add,
Sub,
And,
Or,
}
pub(crate) fn counter_binary_op(l_id: HeapId, r_id: HeapId, op: CounterOp, vm: &mut VM<'_>) -> RunResult<Value> {
let mut result = Dict::new();
result.make_counter();
let result_id = vm.heap.allocate(HeapData::Dict(result));
let mut result_guard = DropGuard::new(Value::Ref(result_id), vm);
let vm = result_guard.ctx();
match op {
CounterOp::Add => {
let l_pairs = counter_snapshot(l_id, vm);
counter_bump_all(result_id, l_pairs, false, vm)?;
let r_pairs = counter_snapshot(r_id, vm);
counter_bump_all(result_id, r_pairs, false, vm)?;
}
CounterOp::Sub => {
let l_pairs = counter_snapshot(l_id, vm);
counter_bump_all(result_id, l_pairs, false, vm)?;
let r_pairs = counter_snapshot(r_id, vm);
counter_bump_all(result_id, r_pairs, true, vm)?;
}
CounterOp::Or => counter_binary_extreme(result_id, l_id, r_id, ExtremeOp::Max, vm)?,
CounterOp::And => counter_binary_extreme(result_id, l_id, r_id, ExtremeOp::Min, vm)?,
}
counter_retain_positive(result_id, vm)?;
Ok(result_guard.into_inner())
}
#[derive(Clone, Copy)]
enum ExtremeOp {
Min,
Max,
}
fn counter_binary_extreme(
result_id: HeapId,
l_id: HeapId,
r_id: HeapId,
op: ExtremeOp,
vm: &mut VM<'_>,
) -> RunResult<()> {
let l_pairs = counter_snapshot(l_id, vm).into_iter();
defer_drop_mut!(l_pairs, vm);
for entry in l_pairs.by_ref() {
let mut entry = DropGuard::new(entry, vm);
let ((key, count), vm) = entry.as_parts_mut();
let other = counter_lookup(r_id, key, vm)?.unwrap_or(Value::Int(0));
let mut other = DropGuard::new(other, vm);
let (other_count, vm) = other.as_parts_mut();
let count_is_smaller = count_holds(count, other_count, CountCmp::Lt, vm)?;
let keep_count = match op {
ExtremeOp::Max => !count_is_smaller,
ExtremeOp::Min => count_is_smaller,
};
let other = other.into_inner();
let ((key, count), vm) = entry.into_parts();
let picked = if keep_count {
other.drop_with(vm);
count
} else {
count.drop_with(vm);
other
};
counter_set(result_id, key, picked, vm)?;
}
if matches!(op, ExtremeOp::Max) {
let r_pairs = counter_snapshot(r_id, vm).into_iter();
defer_drop_mut!(r_pairs, vm);
for entry in r_pairs.by_ref() {
let mut entry = DropGuard::new(entry, vm);
let ((key, _), vm) = entry.as_parts_mut();
if let Some(existing) = counter_lookup(l_id, key, vm)? {
existing.drop_with(vm);
} else {
let ((key, count), vm) = entry.into_parts();
counter_set(result_id, key, count, vm)?;
}
}
}
Ok(())
}
#[derive(Clone, Copy)]
pub(crate) enum CounterCmp {
Le,
Lt,
Ge,
Gt,
}
pub(crate) fn counter_compare(l_id: HeapId, r_id: HeapId, cmp: CounterCmp, vm: &mut VM<'_>) -> RunResult<bool> {
let (op, strict) = match cmp {
CounterCmp::Le => (CountCmp::Le, false),
CounterCmp::Lt => (CountCmp::Le, true),
CounterCmp::Ge => (CountCmp::Ge, false),
CounterCmp::Gt => (CountCmp::Ge, true),
};
let mut holds = true;
let mut equal = true;
let pairs = counter_union_counts(l_id, r_id, vm)?.into_iter();
defer_drop_mut!(pairs, vm);
for pair in pairs.by_ref() {
defer_drop!(pair, vm);
let (l, r) = pair;
match l.py_cmp(r, vm)? {
CmpOrder::Ordered(Ordering::Equal) => {}
CmpOrder::Ordered(ordering) => {
holds &= op.holds(ordering);
equal = false;
}
CmpOrder::Unordered => {
holds = false;
equal = false;
}
CmpOrder::Incomparable => {
return Err(ExcType::type_error_ordering(
op.symbol(),
&l.py_type_name(vm),
&r.py_type_name(vm),
));
}
}
if !holds {
break;
}
}
Ok(holds && (!strict || !equal))
}
fn counter_union_counts(l_id: HeapId, r_id: HeapId, vm: &mut VM<'_>) -> RunResult<Vec<(Value, Value)>> {
let mut pairs = DropGuard::new(Vec::new(), vm);
for (keys_id, other_id, flip) in [(l_id, r_id, false), (r_id, l_id, true)] {
let (collected, vm) = pairs.as_parts_mut();
let entries = counter_snapshot(keys_id, vm).into_iter();
defer_drop_mut!(entries, vm);
for entry in entries.by_ref() {
let mut entry = DropGuard::new(entry, vm);
let ((key, _), vm) = entry.as_parts_mut();
let other = counter_lookup(other_id, key, vm)?;
match (flip, other) {
(true, Some(other)) => other.drop_with(vm),
(true, None) => {
let ((key, count), vm) = entry.into_parts();
key.drop_with(vm);
collected.push((Value::Int(0), count));
}
(false, other) => {
let ((key, count), vm) = entry.into_parts();
key.drop_with(vm);
collected.push((count, other.unwrap_or(Value::Int(0))));
}
}
}
}
Ok(pairs.into_inner())
}
pub(crate) fn counter_inplace_op(l_id: HeapId, rhs: &Value, op: CounterOp, vm: &mut VM<'_>) -> RunResult<()> {
match op {
CounterOp::Add => counter_bump_all(l_id, counter_snapshot(counter_require_mapping(rhs, vm)?, vm), false, vm)?,
CounterOp::Sub => counter_bump_all(l_id, counter_snapshot(counter_require_mapping(rhs, vm)?, vm), true, vm)?,
CounterOp::Or => {
let r_id = counter_require_mapping(rhs, vm)?;
let r_pairs = counter_snapshot(r_id, vm).into_iter();
defer_drop_mut!(r_pairs, vm);
for entry in r_pairs.by_ref() {
let mut entry = DropGuard::new(entry, vm);
let ((key, other_count), vm) = entry.as_parts_mut();
let count = counter_lookup(l_id, key, vm)?.unwrap_or(Value::Int(0));
let bigger = count_holds(other_count, &count, CountCmp::Gt, vm);
count.drop_with(vm);
if bigger? {
let ((key, other_count), vm) = entry.into_parts();
counter_set(l_id, key, other_count, vm)?;
}
}
}
CounterOp::And => {
let pairs = counter_snapshot(l_id, vm).into_iter();
defer_drop_mut!(pairs, vm);
for entry in pairs.by_ref() {
let mut entry = DropGuard::new(entry, vm);
let ((key, count), vm) = entry.as_parts_mut();
let other_count = rhs.py_getitem(key, vm)?;
let mut other_count = DropGuard::new(other_count, vm);
let (other, vm) = other_count.as_parts_mut();
let is_smaller = count_holds(other, count, CountCmp::Lt, vm)?;
if is_smaller {
let smaller = other_count.into_inner();
let ((key, count), vm) = entry.into_parts();
count.drop_with(vm);
counter_set(l_id, key, smaller, vm)?;
}
}
}
}
counter_retain_positive(l_id, vm)
}
pub(crate) fn counter_unary_op(id: HeapId, negate: bool, vm: &mut VM<'_>) -> RunResult<Value> {
let mut result = Dict::new();
result.make_counter();
let result_id = vm.heap.allocate(HeapData::Dict(result));
let mut result_guard = DropGuard::new(Value::Ref(result_id), vm);
let vm = result_guard.ctx();
{
let pairs = counter_snapshot(id, vm).into_iter();
defer_drop_mut!(pairs, vm);
for entry in pairs.by_ref() {
let mut entry = DropGuard::new(entry, vm);
let ((_, count), vm) = entry.as_parts_mut();
let negated = if negate {
if !count_holds(count, &Value::Int(0), CountCmp::Lt, vm)? {
continue;
}
Some(count_arith(&Value::Int(0), count, true, vm)?)
} else {
None
};
let ((key, count), vm) = entry.into_parts();
let stored = match negated {
Some(negated) => {
count.drop_with(vm);
negated
}
None => count,
};
counter_set(result_id, key, stored, vm)?;
}
}
counter_retain_positive(result_id, vm)?;
Ok(result_guard.into_inner())
}
fn counter_bump_all(id: HeapId, pairs: Vec<(Value, Value)>, subtract: bool, vm: &mut VM<'_>) -> RunResult<()> {
let pairs = pairs.into_iter();
defer_drop_mut!(pairs, vm);
for (key, delta) in pairs.by_ref() {
let outcome = counter_bump(id, key, &delta, subtract, false, vm);
delta.drop_with(vm);
outcome?;
}
Ok(())
}
fn counter_snapshot(id: HeapId, vm: &mut VM<'_>) -> Vec<(Value, Value)> {
let HeapReadOutput::Dict(dict) = vm.heap.read(id) else {
unreachable!("counter_snapshot on a non-dict heap entry");
};
let dict = dict.get(vm.heap);
dict.iter()
.map(|(k, v)| (k.clone_with_heap(vm.heap), v.clone_with_heap(vm.heap)))
.collect()
}
fn counter_require_mapping(rhs: &Value, vm: &mut VM<'_>) -> RunResult<HeapId> {
match rhs {
Value::Ref(id) if matches!(vm.heap.get(*id), HeapData::Dict(_)) => Ok(*id),
other => Err(ExcType::attribute_error(other.py_type_name(vm), "items")),
}
}
fn counter_lookup(id: HeapId, key: &Value, vm: &mut VM<'_>) -> RunResult<Option<Value>> {
let HeapReadOutput::Dict(dict) = vm.heap.read(id) else {
unreachable!("counter_lookup on a non-dict heap entry");
};
dict.dict_get(key, vm)
}
fn counter_set(id: HeapId, key: Value, count: Value, vm: &mut VM<'_>) -> RunResult<()> {
let HeapReadOutput::Dict(mut dict) = vm.heap.read(id) else {
unreachable!("counter_set on a non-dict heap entry");
};
if let Some(old) = dict.set(key, count, vm)? {
old.drop_with(vm);
}
Ok(())
}
fn counter_retain_positive(id: HeapId, vm: &mut VM<'_>) -> RunResult<()> {
let mut remove = DropGuard::new(Vec::new(), vm);
{
let (remove, vm) = remove.as_parts_mut();
let entries = counter_snapshot(id, vm).into_iter();
defer_drop_mut!(entries, vm);
for entry in entries.by_ref() {
let mut entry = DropGuard::new(entry, vm);
let ((_, count), vm) = entry.as_parts_mut();
let positive = count_is_positive(count, vm)?;
if !positive {
let ((key, count), vm) = entry.into_parts();
count.drop_with(vm);
remove.push(key);
}
}
}
let (remove, vm) = remove.into_parts();
let remove = remove.into_iter();
defer_drop_mut!(remove, vm);
for key in remove.by_ref() {
defer_drop!(key, vm);
let HeapReadOutput::Dict(mut dict) = vm.heap.read(id) else {
unreachable!("counter_retain_positive on a non-dict heap entry");
};
if let Some(old) = dict.pop(key, vm)? {
old.drop_with(vm);
}
}
Ok(())
}