use super::*;
use crate::Write;
use crate::exec_state::Internal;
use inner::MultiSet;
#[derive(Clone, Debug, PartialEq, Eq, Hash)]
pub struct MultiSetContainer {
pub do_rebuild: bool,
pub data: MultiSet<Value>,
}
fn normalize_multiset_term(termdag: &mut TermDag, mut children: Vec<TermId>) -> TermId {
termdag.sort_terms_by_ast(&mut children);
termdag.app("multiset-of".into(), children)
}
fn multiset_term_children(termdag: &TermDag, term: TermId) -> Option<Vec<TermId>> {
match termdag.get(term) {
Term::App(head, children) if head == "multiset-of" => Some(children.clone()),
_ => None,
}
}
impl ContainerValue for MultiSetContainer {
fn rebuild_contents(&mut self, rebuilder: &dyn ValueRebuilder) -> bool {
if self.do_rebuild {
let mut xs: Vec<_> = self.data.iter().copied().collect();
let changed = rebuilder.rebuild_slice(&mut xs);
self.data = xs.into_iter().collect();
changed
} else {
false
}
}
fn iter(&self) -> impl Iterator<Item = Value> + '_ {
self.data.iter().copied()
}
}
#[derive(Clone, Debug)]
pub struct MultiSetSort {
name: String,
element: ArcSort,
}
impl MultiSetSort {
pub fn element(&self) -> ArcSort {
self.element.clone()
}
}
impl Presort for MultiSetSort {
fn presort_name() -> &'static str {
"MultiSet"
}
fn reserved_primitives() -> Vec<&'static str> {
vec![
"multiset-of",
"multiset-single",
"multiset-insert",
"multiset-remove",
"multiset-remove-swapped",
"multiset-subtract",
"multiset-subtract-swapped",
"multiset-length",
"multiset-contains",
"multiset-not-contains",
"multiset-contains-swapped",
"multiset-not-contains-swapped",
"multiset-intersection",
"multiset-sum",
"multiset-reset-counts",
"multiset-pick-max",
"multiset-count",
"multiset-sum-multisets",
"unstable-multiset-map",
"unstable-multiset-filter",
"unstable-multiset-filter-not",
"unstable-multiset-reduce",
"unstable-multiset-fill-index",
"unstable-multiset-clear-index",
"unstable-multiset-flat-map",
]
}
fn make_sort(
typeinfo: &mut TypeInfo,
name: String,
args: &[Expr],
span: Span,
) -> Result<ArcSort, TypeError> {
if let [Expr::Var(arg_span, e)] = args {
let e = typeinfo
.get_sort_by_name(e)
.ok_or(TypeError::UndefinedSort(e.clone(), arg_span.clone()))?;
let out = Self {
name,
element: e.clone(),
};
Ok(out.to_arcsort())
} else {
Err(TypeError::BadPresortArguments(
Self::presort_name().to_owned(),
span,
))
}
}
}
impl ContainerSort for MultiSetSort {
type Container = MultiSetContainer;
fn name(&self) -> &str {
&self.name
}
fn inner_sorts(&self) -> Vec<ArcSort> {
vec![self.element.clone()]
}
fn is_eq_container_sort(&self) -> bool {
self.element.is_eq_sort() || self.element.is_eq_container_sort()
}
fn inner_values(
&self,
container_values: &ContainerValues,
value: Value,
) -> Vec<(ArcSort, Value)> {
let val = container_values
.get_val::<MultiSetContainer>(value)
.unwrap()
.clone();
val.data
.iter()
.map(|k| (self.element.clone(), *k))
.collect()
}
fn register_primitives(&self, eg: &mut EGraph) {
let arc = self.clone().to_arcsort();
let multiset_of_validator = |termdag: &mut TermDag, args: &[TermId]| -> Option<TermId> {
Some(normalize_multiset_term(termdag, args.to_vec()))
};
let multiset_length_validator =
|termdag: &mut TermDag, args: &[TermId]| -> Option<TermId> {
let [ms] = args else { return None };
let len = multiset_term_children(termdag, *ms)?.len() as i64;
Some(termdag.lit(Literal::Int(len)))
};
let multiset_contains_validator =
|termdag: &mut TermDag, args: &[TermId]| -> Option<TermId> {
let [ms, value] = args else { return None };
multiset_term_children(termdag, *ms)?
.contains(value)
.then(|| termdag.lit(Literal::Unit))
};
let multiset_not_contains_validator =
|termdag: &mut TermDag, args: &[TermId]| -> Option<TermId> {
let [ms, value] = args else { return None };
let contains = multiset_term_children(termdag, *ms)?.contains(value);
(!contains).then(|| termdag.lit(Literal::Unit))
};
add_primitive_with_validator!(eg, "multiset-of" = {self.clone(): MultiSetSort} [xs: # (self.element())] -> @MultiSetContainer (arc) { MultiSetContainer {
do_rebuild: self.ctx.is_eq_container_sort(),
data: xs.collect()
} }, multiset_of_validator);
add_primitive!(eg, "multiset-single" = {self.clone(): MultiSetSort} |x: # (self.element()), i: i64| -?> @MultiSetContainer (arc) {
i.try_into().ok().map(|i|
MultiSetContainer {
do_rebuild: self.ctx.is_eq_container_sort(),
data: std::iter::repeat_n(x, i).collect()
})
});
add_primitive!(eg, "multiset-pick" = |xs: @MultiSetContainer (arc)| -?> # (self.element()) { xs.data.pick().copied() });
add_primitive!(eg, "multiset-insert" = |mut xs: @MultiSetContainer (arc), x: # (self.element())| -> @MultiSetContainer (arc) { MultiSetContainer { data: xs.data.insert( x) , ..xs } });
add_primitive!(eg, "multiset-remove" = |mut xs: @MultiSetContainer (arc), x: # (self.element())| -?> @MultiSetContainer (arc) { Some(MultiSetContainer { data: xs.data.remove(&x)?, ..xs } )});
add_primitive!(eg, "multiset-remove-swapped" = |x: # (self.element()), mut xs: @MultiSetContainer (arc)| -?> @MultiSetContainer (arc) { Some(MultiSetContainer { data: xs.data.remove(&x)?, ..xs }) });
add_primitive!(eg, "multiset-subtract" = |mut xs: @MultiSetContainer (arc), other: @MultiSetContainer (arc)| -?> @MultiSetContainer (arc) { Some(MultiSetContainer { data: xs.data.subtract(&other.data)?, ..xs }) });
add_primitive!(eg, "multiset-subtract-swapped" = |other: @MultiSetContainer (arc), mut xs: @MultiSetContainer (arc)| -?> @MultiSetContainer (arc) { Some(MultiSetContainer { data: xs.data.subtract(&other.data)?, ..xs }) });
add_primitive_with_validator!(eg, "multiset-length" = |xs: @MultiSetContainer (arc)| -> i64 { xs.data.len() as i64 }, multiset_length_validator);
add_primitive_with_validator!(eg, "multiset-contains" = |xs: @MultiSetContainer (arc), x: # (self.element())| -?> () { ( xs.data.contains(&x)).then_some(()) }, multiset_contains_validator);
add_primitive_with_validator!(eg, "multiset-not-contains" = |xs: @MultiSetContainer (arc), x: # (self.element())| -?> () { (!xs.data.contains(&x)).then_some(()) }, multiset_not_contains_validator);
add_primitive!(eg, "multiset-contains-swapped" = |x: # (self.element()), xs: @MultiSetContainer (arc)| -?> () { (xs.data.contains(&x)).then_some(()) });
add_primitive!(eg, "multiset-not-contains-swapped" = |x: # (self.element()), xs: @MultiSetContainer (arc)| -?> () { (!xs.data.contains(&x)).then_some(()) });
add_primitive!(eg, "multiset-intersection" = |xs: @MultiSetContainer (arc), ys: @MultiSetContainer (arc)| -> @MultiSetContainer (arc) { MultiSetContainer { data: xs.data.intersection(ys.data), ..xs } });
add_primitive!(eg, "multiset-sum" = |xs: @MultiSetContainer (arc), ys: @MultiSetContainer (arc)| -> @MultiSetContainer (arc) { MultiSetContainer { data: xs.data.sum(ys.data), ..xs } });
add_primitive!(eg, "multiset-reset-counts" = |mut xs: @MultiSetContainer (arc)| -> @MultiSetContainer (arc) { {
let mut new_data = MultiSet::<Value>::new();
for (v, _) in xs.data.iter_counts() {
new_data.insert_multiple_mut(v, 1);
}
MultiSetContainer { data: new_data, ..xs }
}});
add_primitive!(eg, "multiset-pick-max" = |xs: @MultiSetContainer (arc)| -?> # (self.element()) {
Some(xs.data.iter_counts().max_by_key(|(_, c)| *c)?.0)
});
add_primitive!(eg, "multiset-count" = |xs: @MultiSetContainer (arc), x: # (self.element())| -> i64 {
xs.data.iter_counts().find(|(v, _)| *v == x).map(|(_, c)| c as i64).unwrap_or(0)
});
for other_multiset_sort in eg.type_info.get_arcsorts_by(|f| {
f.name() == self.element.name()
&& f.value_type() == Some(TypeId::of::<MultiSetContainer>())
}) {
eg.add_pure_primitive(
SumMultisets {
name: "multiset-sum-multisets".into(),
multiset: other_multiset_sort.clone(),
multiset_of_multisets: arc.clone(),
},
None,
);
}
let all_ms_sorts = eg
.type_info
.get_arcsorts_by(|f| f.value_type() == Some(TypeId::of::<MultiSetContainer>()));
for fn_sort in eg.type_info.get_sorts::<FunctionSort>() {
for ms_sort in &all_ms_sorts {
try_registering_multiset_map(eg, fn_sort.clone(), ms_sort.clone(), arc.clone());
if ms_sort.name() != arc.name() {
try_registering_multiset_map(eg, fn_sort.clone(), arc.clone(), ms_sort.clone());
}
}
try_registering_multiset_non_map_primitives(eg, fn_sort.clone(), arc.clone());
}
if self.element.is_eq_sort() {
eg.add_write_primitive(
UnionValues {
name: "multiset-union-values".into(),
multiset: arc.clone(),
element: self.element.clone(),
},
None,
);
}
}
fn reconstruct_termdag(
&self,
_container_values: &ContainerValues,
_value: Value,
termdag: &mut TermDag,
element_terms: Vec<TermId>,
) -> TermId {
normalize_multiset_term(termdag, element_terms)
}
fn rebuild_container_normalizer(&self) -> Option<(String, PrimitiveValidator)> {
Some((
"multiset-of".to_owned(),
Arc::new(|termdag: &mut TermDag, args: &[TermId]| {
Some(normalize_multiset_term(termdag, args.to_vec()))
}),
))
}
fn serialized_name(&self, _container_values: &ContainerValues, _: Value) -> String {
"multiset-of".to_owned()
}
}
pub(crate) fn try_registering_multiset_map(
eg: &mut EGraph,
fn_: Arc<FunctionSort>,
input_ms: ArcSort,
output_ms: ArcSort,
) {
if fn_.inputs().len() != 1
|| fn_.inputs()[0].name() != input_ms.inner_sorts()[0].name()
|| fn_.output().name() != output_ms.inner_sorts()[0].name()
{
return;
}
eg.add_pure_primitive(
Map {
name: "unstable-multiset-map".into(),
multiset: input_ms,
output_multiset: output_ms,
fn_: fn_.clone(),
},
None,
);
}
pub(crate) fn register_multiset_primitives_for_function(eg: &mut EGraph, fn_: Arc<FunctionSort>) {
let all_ms_sorts = eg
.type_info
.get_arcsorts_by(|f| f.value_type() == Some(TypeId::of::<MultiSetContainer>()));
for input_ms in &all_ms_sorts {
for output_ms in &all_ms_sorts {
try_registering_multiset_map(eg, fn_.clone(), input_ms.clone(), output_ms.clone());
}
}
for ms_sort in &all_ms_sorts {
try_registering_multiset_non_map_primitives(eg, fn_.clone(), ms_sort.clone());
}
}
fn try_registering_multiset_non_map_primitives(
eg: &mut EGraph,
fn_: Arc<FunctionSort>,
multiset: ArcSort,
) {
let element = multiset.inner_sorts()[0].clone();
let element_name = element.name();
if fn_.inputs().len() == 1
&& fn_.inputs()[0].name() == element_name
&& fn_.output().name() == "Unit"
{
eg.add_pure_primitive(
Filter {
name: "unstable-multiset-filter".into(),
multiset: multiset.clone(),
fn_: fn_.clone(),
skip_empty: true,
},
None,
);
eg.add_pure_primitive(
Filter {
name: "unstable-multiset-filter-not".into(),
multiset: multiset.clone(),
fn_: fn_.clone(),
skip_empty: false,
},
None,
);
}
if fn_.inputs().len() == 2
&& fn_.inputs()[0].name() == element_name
&& fn_.inputs()[1].name() == element_name
&& fn_.output().name() == element_name
{
eg.add_pure_primitive(
Reduce {
name: "unstable-multiset-reduce".into(),
multiset: multiset.clone(),
fn_: fn_.clone(),
element: element.clone(),
},
None,
);
}
if fn_.inputs().len() == 2
&& fn_.inputs()[0].name() == multiset.name()
&& fn_.inputs()[1].name() == element_name
&& fn_.output().name() == "i64"
{
let unit = eg.type_info.get_sort_by_name("Unit").unwrap().clone();
eg.add_full_primitive(
FillIndex {
name: "unstable-multiset-fill-index".into(),
multiset: multiset.clone(),
unit: unit.clone(),
fn_: fn_.clone(),
},
None,
);
eg.add_write_primitive(
ClearIndex {
name: "unstable-multiset-clear-index".into(),
multiset: multiset.clone(),
unit,
fn_: fn_.clone(),
},
None,
);
}
if fn_.inputs().len() == 1
&& fn_.inputs()[0].name() == element_name
&& fn_.output().name() == multiset.name()
{
eg.add_pure_primitive(
FlatMap {
name: "unstable-multiset-flat-map".into(),
multiset,
fn_: fn_.clone(),
},
None,
);
}
}
#[derive(Clone)]
struct Map {
name: String,
multiset: ArcSort,
fn_: Arc<FunctionSort>,
output_multiset: ArcSort,
}
impl Primitive for Map {
fn name(&self) -> &str {
&self.name
}
fn get_type_constraints(&self, span: &Span) -> Box<dyn TypeConstraint> {
SimpleTypeConstraint::new(
&self.name,
vec![
self.fn_.clone(),
self.multiset.clone(),
self.output_multiset.clone(),
],
span.clone(),
)
.into_box()
}
}
impl PurePrim for Map {
fn apply<'a, 'db>(
&self,
mut state: crate::PureState<'a, 'db>,
args: &[Value],
) -> Option<Value> {
let fc = state
.container_values()
.get_val::<FunctionContainer>(args[0])
.unwrap()
.clone();
let multiset = state
.container_values()
.get_val::<MultiSetContainer>(args[1])
.unwrap()
.clone();
let mut new_data = MultiSet::<Value>::new();
for (v, c) in multiset.data.iter_counts() {
if let Some(mapped) = state.apply_function(&fc, &[v]) {
new_data.insert_multiple_mut(mapped, c);
}
}
let new_ms = MultiSetContainer {
data: new_data,
..multiset
};
Some(state.register_container(new_ms))
}
}
#[derive(Clone)]
struct FillIndex {
name: String,
multiset: ArcSort,
unit: ArcSort,
fn_: Arc<FunctionSort>,
}
impl Primitive for FillIndex {
fn name(&self) -> &str {
&self.name
}
fn get_type_constraints(&self, span: &Span) -> Box<dyn TypeConstraint> {
SimpleTypeConstraint::new(
self.name(),
vec![self.multiset.clone(), self.fn_.clone(), self.unit.clone()],
span.clone(),
)
.into_box()
}
}
impl FullPrim for FillIndex {
fn apply<'a, 'db>(
&self,
mut state: crate::FullState<'a, 'db>,
args: &[Value],
) -> Option<Value> {
let fc = state
.container_values()
.get_val::<FunctionContainer>(args[1])
.unwrap()
.clone();
let multiset = state
.container_values()
.get_val::<MultiSetContainer>(args[0])
.unwrap()
.clone();
let action = match fc.0 {
ResolvedFunctionId::Constructor(a) | ResolvedFunctionId::Function(a) => a,
ResolvedFunctionId::Primitive { .. } => return None,
};
let unit_val = state.base_values().get::<()>(());
let es = state.raw_exec_state();
for (v, c) in multiset.data.iter_counts() {
let mut row = vec![args[0], v];
if action.lookup(es, &row).is_some() {
break;
}
row.push(es.base_values().get::<i64>(c.try_into().ok()?));
action.insert(es, row.into_iter());
}
Some(unit_val)
}
}
#[derive(Clone)]
struct ClearIndex {
name: String,
multiset: ArcSort,
unit: ArcSort,
fn_: Arc<FunctionSort>,
}
impl Primitive for ClearIndex {
fn name(&self) -> &str {
&self.name
}
fn get_type_constraints(&self, span: &Span) -> Box<dyn TypeConstraint> {
SimpleTypeConstraint::new(
self.name(),
vec![self.multiset.clone(), self.fn_.clone(), self.unit.clone()],
span.clone(),
)
.into_box()
}
}
impl WritePrim for ClearIndex {
fn apply<'a, 'db>(
&self,
mut state: crate::WriteState<'a, 'db>,
args: &[Value],
) -> Option<Value> {
let fc = state
.container_values()
.get_val::<FunctionContainer>(args[1])
.unwrap()
.clone();
let multiset = state
.container_values()
.get_val::<MultiSetContainer>(args[0])
.unwrap()
.clone();
let action = match fc.0 {
ResolvedFunctionId::Constructor(a) | ResolvedFunctionId::Function(a) => a,
ResolvedFunctionId::Primitive { .. } => return None,
};
let unit_val = state.base_values().get::<()>(());
let es = state.raw_exec_state();
for (v, _) in multiset.data.iter_counts() {
action.remove(es, &[args[0], v]);
}
Some(unit_val)
}
}
#[derive(Clone)]
struct FlatMap {
name: String,
multiset: ArcSort,
fn_: Arc<FunctionSort>,
}
impl Primitive for FlatMap {
fn name(&self) -> &str {
&self.name
}
fn get_type_constraints(&self, span: &Span) -> Box<dyn TypeConstraint> {
SimpleTypeConstraint::new(
&self.name,
vec![
self.fn_.clone(),
self.multiset.clone(),
self.multiset.clone(),
],
span.clone(),
)
.into_box()
}
}
impl PurePrim for FlatMap {
fn apply<'a, 'db>(
&self,
mut state: crate::PureState<'a, 'db>,
args: &[Value],
) -> Option<Value> {
let fc = state
.container_values()
.get_val::<FunctionContainer>(args[0])
.unwrap()
.clone();
let multiset = state
.container_values()
.get_val::<MultiSetContainer>(args[1])
.unwrap()
.clone();
let mut new_data = MultiSet::<Value>::new();
for (v, c) in multiset.data.iter_counts() {
let mapped = state.apply_function(&fc, &[v]);
if let Some(mapped_ms) = mapped {
let mapped_ms = state
.container_values()
.get_val::<MultiSetContainer>(mapped_ms)
.unwrap();
for (mv, mc) in mapped_ms.data.iter_counts() {
new_data.insert_multiple_mut(mv, c.checked_mul(mc)?);
}
} else {
new_data.insert_multiple_mut(v, c);
}
}
let new_container = MultiSetContainer {
data: new_data,
..multiset
};
Some(state.register_container(new_container))
}
}
#[derive(Clone)]
struct Filter {
name: String,
multiset: ArcSort,
fn_: Arc<FunctionSort>,
skip_empty: bool,
}
impl Primitive for Filter {
fn name(&self) -> &str {
&self.name
}
fn get_type_constraints(&self, span: &Span) -> Box<dyn TypeConstraint> {
SimpleTypeConstraint::new(
&self.name,
vec![
self.fn_.clone(),
self.multiset.clone(),
self.multiset.clone(),
],
span.clone(),
)
.into_box()
}
}
impl PurePrim for Filter {
fn apply<'a, 'db>(
&self,
mut state: crate::PureState<'a, 'db>,
args: &[Value],
) -> Option<Value> {
let fc = state
.container_values()
.get_val::<FunctionContainer>(args[0])
.unwrap()
.clone();
let multiset = state
.container_values()
.get_val::<MultiSetContainer>(args[1])
.unwrap()
.clone();
let mut new_data = MultiSet::<Value>::new();
for (v, c) in multiset.data.iter_counts() {
let mapped = state.apply_function(&fc, &[v]);
if mapped.is_some() == self.skip_empty {
new_data.insert_multiple_mut(v, c);
}
}
let new_ms = MultiSetContainer {
data: new_data,
..multiset
};
Some(state.register_container(new_ms))
}
}
#[derive(Clone)]
struct SumMultisets {
name: String,
multiset: ArcSort,
multiset_of_multisets: ArcSort,
}
impl Primitive for SumMultisets {
fn name(&self) -> &str {
&self.name
}
fn get_type_constraints(&self, span: &Span) -> Box<dyn TypeConstraint> {
SimpleTypeConstraint::new(
self.name(),
vec![self.multiset_of_multisets.clone(), self.multiset.clone()],
span.clone(),
)
.into_box()
}
}
impl PurePrim for SumMultisets {
fn apply<'a, 'db>(
&self,
mut state: crate::PureState<'a, 'db>,
args: &[Value],
) -> Option<Value> {
let mut data = MultiSet::<Value>::new();
let ms_of_ms = state
.container_values()
.get_val::<MultiSetContainer>(args[0])
.unwrap()
.clone();
for (ms_value, counts) in ms_of_ms.data.iter_counts() {
let ms = state
.container_values()
.get_val::<MultiSetContainer>(ms_value)
.unwrap();
for (v, c) in ms.data.iter_counts() {
data.insert_multiple_mut(v, c.checked_mul(counts)?);
}
}
let multiset = MultiSetContainer {
data,
do_rebuild: self.multiset.is_eq_container_sort(),
};
Some(state.register_container(multiset))
}
}
#[derive(Clone)]
struct Reduce {
name: String,
multiset: ArcSort,
fn_: Arc<FunctionSort>,
element: ArcSort,
}
impl Primitive for Reduce {
fn name(&self) -> &str {
&self.name
}
fn get_type_constraints(&self, span: &Span) -> Box<dyn TypeConstraint> {
SimpleTypeConstraint::new(
&self.name,
vec![
self.fn_.clone(),
self.element.clone(),
self.multiset.clone(),
self.element.clone(),
],
span.clone(),
)
.into_box()
}
}
impl PurePrim for Reduce {
fn apply<'a, 'db>(
&self,
mut state: crate::PureState<'a, 'db>,
args: &[Value],
) -> Option<Value> {
let fc = state
.container_values()
.get_val::<FunctionContainer>(args[0])
.unwrap()
.clone();
let initial = args[1];
let multiset = state
.container_values()
.get_val::<MultiSetContainer>(args[2])
.unwrap()
.clone();
let mut values = multiset.data.iter().cloned().collect::<Vec<_>>();
let mut acc = if values.is_empty() {
initial
} else {
values.remove(0)
};
for v in values {
acc = state.apply_function(&fc, &[acc, v])?;
}
Some(acc)
}
}
#[derive(Clone)]
struct UnionValues {
name: String,
multiset: ArcSort,
element: ArcSort,
}
impl Primitive for UnionValues {
fn name(&self) -> &str {
&self.name
}
fn get_type_constraints(&self, span: &Span) -> Box<dyn TypeConstraint> {
SimpleTypeConstraint::new(
self.name(),
vec![self.multiset.clone(), self.element.clone()],
span.clone(),
)
.into_box()
}
}
impl WritePrim for UnionValues {
fn apply<'a, 'db>(
&self,
mut state: crate::WriteState<'a, 'db>,
args: &[Value],
) -> Option<Value> {
let values = state
.container_values()
.get_val::<MultiSetContainer>(args[0])?
.clone()
.data;
let values: Vec<_> = values.iter_counts().map(|(v, _c)| v).collect();
if values.is_empty() {
return None;
}
let first = values[0];
for v in values.into_iter().skip(1) {
state.union(first, v).ok()?;
}
Some(first)
}
}
mod inner {
use std::collections::BTreeMap;
use std::hash::Hash;
#[derive(Debug, Default, Hash, Eq, PartialEq, Clone)]
pub struct MultiSet<T: Clone + Hash + Ord>(
BTreeMap<T, usize>,
usize,
);
impl<T: Clone + Hash + Ord> MultiSet<T> {
pub fn new() -> Self {
MultiSet(BTreeMap::new(), 0)
}
pub fn contains(&self, value: &T) -> bool {
self.0.contains_key(value)
}
pub fn len(&self) -> usize {
self.1
}
pub fn iter(&self) -> impl Iterator<Item = &T> {
self.0.iter().flat_map(|(k, v)| std::iter::repeat_n(k, *v))
}
pub fn iter_counts(&self) -> impl Iterator<Item = (T, usize)> {
self.0.iter().map(|(k, v)| (k.clone(), *v))
}
pub fn pick(&self) -> Option<&T> {
self.0.keys().next()
}
pub fn insert(mut self, value: T) -> MultiSet<T> {
self.insert_multiple_mut(value, 1);
self
}
pub fn remove(mut self, value: &T) -> Option<MultiSet<T>> {
if let Some(v) = self.0.get(value) {
self.1 -= 1;
if *v == 1 {
self.0.remove(value);
} else {
self.0.insert(value.clone(), v - 1);
}
Some(self)
} else {
None
}
}
pub fn subtract(mut self, other: &MultiSet<T>) -> Option<MultiSet<T>> {
for (k, v) in other.0.iter() {
if let Some(self_v) = self.0.get_mut(k) {
if *self_v < *v {
return None;
}
*self_v -= *v;
self.1 -= *v;
if *self_v == 0 {
self.0.remove(k);
}
} else {
return None;
}
}
Some(self)
}
pub fn insert_multiple_mut(&mut self, value: T, n: usize) {
self.1 += n;
if let Some(v) = self.0.get(&value) {
self.0.insert(value, v + n);
} else {
self.0.insert(value, n);
}
}
pub fn sum(mut self, MultiSet(other_map, other_count): Self) -> Self {
let target_count = self.1 + other_count;
for (k, v) in other_map {
self.insert_multiple_mut(k, v);
}
assert_eq!(self.1, target_count);
self
}
pub fn intersection(self, MultiSet(other_map, _): Self) -> Self {
let mut new_map = BTreeMap::new();
for (k, v) in self.0.into_iter() {
if let Some(other_v) = other_map.get(&k) {
let new_v = std::cmp::min(v, *other_v);
new_map.insert(k, new_v);
}
}
let new_count = new_map.values().sum();
MultiSet(new_map, new_count)
}
}
impl<T: Clone + Hash + Ord> FromIterator<T> for MultiSet<T> {
fn from_iter<I: IntoIterator<Item = T>>(iter: I) -> Self {
let mut multiset = MultiSet::new();
for value in iter {
multiset.insert_multiple_mut(value, 1);
}
multiset
}
}
}