use std::{collections::HashMap, sync::Arc};
use hugr_core::builder::{
BuildError, ConditionalBuilder, DFGBuilder, Dataflow, DataflowHugr, DataflowSubContainer,
HugrBuilder, inout_sig,
};
use hugr_core::extension::{SignatureError, TypeDef};
use hugr_core::std_extensions::collections::array::array_type_def;
use hugr_core::std_extensions::collections::borrow_array::borrow_array_type_def;
use hugr_core::types::{CustomType, Signature, Type, TypeArg, TypeEnum, TypeRow};
use hugr_core::{HugrView, IncomingPort, Node, Wire, hugr::hugrmut::HugrMut, ops::Tag};
use itertools::Itertools;
use super::handlers::{copy_discard_array, copy_discard_borrow_array};
use super::{NodeTemplate, ParametricType};
#[deprecated(
note = "`hugr-passes` is deprecated. Use tket::passes instead",
since = "0.26.2"
)]
pub trait Linearizer {
fn insert_copy_discard(
&self,
hugr: &mut impl HugrMut<Node = Node>,
src: Wire,
targets: &[(Node, IncomingPort)],
) -> Result<(), LinearizeError> {
let (tgt_node, tgt_inport) = if targets.len() == 1 {
*targets.first().unwrap()
} else {
let src_parent = hugr
.get_parent(src.node())
.expect("Root node cannot have out edges");
if let Some((tgt, tgt_parent)) = targets.iter().find_map(|(tgt, _)| {
let tgt_parent = hugr
.get_parent(*tgt)
.expect("Root node cannot have incoming edges");
(tgt_parent != src_parent).then_some((*tgt, tgt_parent))
}) {
return Err(LinearizeError::NoLinearNonLocalEdges {
src: src.node(),
src_parent,
tgt,
tgt_parent,
});
}
let sig = hugr.signature(src.node()).unwrap();
let typ = sig.port_type(src.source()).unwrap().clone();
let copy_discard_op = self
.copy_discard_op(&typ, targets.len())?
.add_hugr(hugr, src_parent)
.map_err(|e| LinearizeError::NestedTemplateError(Box::new(typ), Box::new(e)))?;
for (n, (tgt_node, tgt_port)) in targets.iter().enumerate() {
hugr.connect(copy_discard_op, n, *tgt_node, *tgt_port);
}
(copy_discard_op, 0.into())
};
hugr.connect(src.node(), src.source(), tgt_node, tgt_inport);
Ok(())
}
fn copy_discard_op(
&self,
typ: &Type,
num_outports: usize,
) -> Result<NodeTemplate, LinearizeError>;
}
#[derive(Clone)]
#[deprecated(
note = "`hugr-passes` is deprecated. Use tket::passes instead",
since = "0.26.2"
)]
pub struct DelegatingLinearizer {
copy_discard: HashMap<CustomType, (NodeTemplate, NodeTemplate)>,
copy_discard_parametric: HashMap<
ParametricType,
Arc<
dyn Fn(&[TypeArg], usize, &CallbackHandler<'_>) -> Result<NodeTemplate, LinearizeError>,
>,
>,
}
impl Default for DelegatingLinearizer {
fn default() -> Self {
let mut res = Self::new_empty();
res.register_callback(array_type_def(), copy_discard_array);
res.register_callback(borrow_array_type_def(), copy_discard_borrow_array);
res
}
}
#[deprecated(
note = "`hugr-passes` is deprecated. Use tket::passes instead",
since = "0.26.2"
)]
pub struct CallbackHandler<'a>(&'a DelegatingLinearizer);
#[derive(Clone, Debug, thiserror::Error, PartialEq)]
#[expect(missing_docs)]
#[non_exhaustive]
#[deprecated(
note = "`hugr-passes` is deprecated. Use tket::passes instead",
since = "0.26.2"
)]
pub enum LinearizeError {
#[error("Need copy/discard op for {_0}")]
NeedCopyDiscard(Box<Type>),
#[error("Copy/discard op for {typ} with {num_outports} outputs had wrong signature {sig:?}")]
WrongSignature {
typ: Box<Type>,
num_outports: usize,
sig: Option<Box<Signature>>,
},
#[error(
"Cannot add nonlocal edge for linear type from {src} (with parent {src_parent}) to {tgt} (with parent {tgt_parent}).
Try using LocalizeEdges pass first."
)]
NoLinearNonLocalEdges {
src: Node,
src_parent: Node,
tgt: Node,
tgt_parent: Node,
},
#[error(transparent)]
SignatureError(#[from] SignatureError),
#[error("Cannot linearize type {_0}")]
UnsupportedType(Box<Type>),
#[error("Type {_0} is copyable")]
CopyableType(Box<Type>),
#[error("Could not generate NodeTemplate for contained type {0} because {1}")]
NestedTemplateError(Box<Type>, Box<BuildError>),
}
impl DelegatingLinearizer {
#[must_use]
pub fn new_empty() -> Self {
Self {
copy_discard: Default::default(),
copy_discard_parametric: Default::default(),
}
}
pub fn register_simple(
&mut self,
cty: CustomType,
copy: NodeTemplate,
discard: NodeTemplate,
) -> Result<(), LinearizeError> {
let typ = Type::new_extension(cty.clone());
if typ.copyable() {
return Err(LinearizeError::CopyableType(Box::new(typ)));
}
check_sig(©, &typ, 2)?;
check_sig(&discard, &typ, 0)?;
self.copy_discard.insert(cty, (copy, discard));
Ok(())
}
pub fn register_callback(
&mut self,
src: &TypeDef,
copy_discard_fn: impl Fn(
&[TypeArg],
usize,
&CallbackHandler<'_>,
) -> Result<NodeTemplate, LinearizeError>
+ 'static,
) {
self.copy_discard_parametric
.insert(src.into(), Arc::new(copy_discard_fn));
}
}
fn check_sig(tmpl: &NodeTemplate, typ: &Type, num_outports: usize) -> Result<(), LinearizeError> {
tmpl.check_signature(
&[typ.clone()].into(),
&vec![typ.clone(); num_outports].into(),
)
.map_err(|sig| LinearizeError::WrongSignature {
typ: Box::new(typ.clone()),
num_outports,
sig: sig.map(Box::new),
})
}
impl Linearizer for DelegatingLinearizer {
fn copy_discard_op(
&self,
typ: &Type,
num_outports: usize,
) -> Result<NodeTemplate, LinearizeError> {
if typ.copyable() {
return Err(LinearizeError::CopyableType(Box::new(typ.clone())));
}
assert!(num_outports != 1);
match typ.as_type_enum() {
TypeEnum::Sum(sum_type) => {
let variants = sum_type
.variants()
.map(|trv| trv.clone().try_into())
.collect::<Result<Vec<TypeRow>, _>>()?;
let mut cb = ConditionalBuilder::new(
variants.clone(),
vec![],
vec![sum_type.clone().into(); num_outports],
)
.unwrap();
for (tag, variant) in variants.iter().enumerate() {
let mut case_b = cb.case_builder(tag).unwrap();
let mut elems_for_copy = vec![vec![]; num_outports];
for (inp, ty) in case_b.input_wires().zip_eq(variant.iter()) {
let inp_copies = if ty.copyable() {
std::iter::repeat_n(inp, num_outports).collect::<Vec<_>>()
} else {
self.copy_discard_op(ty, num_outports)?
.add(&mut case_b, [inp])
.unwrap()
.outputs()
.collect()
};
for (src, elems) in inp_copies.into_iter().zip_eq(elems_for_copy.iter_mut())
{
elems.push(src);
}
}
let t = Tag::new(tag, variants.clone());
let outputs = elems_for_copy
.into_iter()
.map(|elems| {
let [copy] = case_b
.add_dataflow_op(t.clone(), elems)
.unwrap()
.outputs_arr();
copy
})
.collect::<Vec<_>>(); case_b.finish_with_outputs(outputs).unwrap();
}
Ok(NodeTemplate::CompoundOp(Box::new(
cb.finish_hugr().unwrap(),
)))
}
TypeEnum::Extension(cty) => {
if let Some((copy, discard)) = self.copy_discard.get(cty) {
Ok(if num_outports == 0 {
discard.clone()
} else {
let mut dfb = DFGBuilder::new(inout_sig(
[typ.clone()],
vec![typ.clone(); num_outports],
))
.unwrap();
let [mut src] = dfb.input_wires_arr();
let mut outputs = vec![];
for _ in 0..num_outports - 1 {
let [out0, out1] =
copy.clone().add(&mut dfb, [src]).unwrap().outputs_arr();
outputs.push(out0);
src = out1;
}
outputs.push(src);
NodeTemplate::CompoundOp(Box::new(
dfb.finish_hugr_with_outputs(outputs).unwrap(),
))
})
} else {
let copy_discard_fn = self
.copy_discard_parametric
.get(&cty.into())
.ok_or_else(|| LinearizeError::NeedCopyDiscard(Box::new(typ.clone())))?;
let tmpl = copy_discard_fn(cty.args(), num_outports, &CallbackHandler(self))?;
check_sig(&tmpl, typ, num_outports)?;
Ok(tmpl)
}
}
TypeEnum::Function(_) => panic!("Ruled out above as copyable"),
_ => Err(LinearizeError::UnsupportedType(Box::new(typ.clone()))),
}
}
}
impl Linearizer for CallbackHandler<'_> {
fn copy_discard_op(
&self,
typ: &Type,
num_outports: usize,
) -> Result<NodeTemplate, LinearizeError> {
self.0.copy_discard_op(typ, num_outports)
}
}
#[cfg(test)]
mod test {
use std::collections::HashMap;
use std::sync::Arc;
use hugr_core::builder::{
Container, DFGBuilder, Dataflow, DataflowHugr, DataflowSubContainer, HugrBuilder, inout_sig,
};
use hugr_core::Visibility;
use hugr_core::extension::prelude::{option_type, qb_t, usize_t};
use hugr_core::extension::{
CustomSignatureFunc, OpDef, SignatureError, SignatureFunc, TypeDefBound, Version,
};
use hugr_core::hugr::ValidationError;
use hugr_core::hugr::hugrmut::HugrMut;
use hugr_core::ops::handle::NodeHandle;
use hugr_core::ops::{DataflowOpTrait, ExtensionOp, OpName, OpType};
use hugr_core::std_extensions::arithmetic::int_types::INT_TYPES;
use hugr_core::std_extensions::collections::array::array_type;
use hugr_core::std_extensions::collections::borrow_array::{BArrayOpDef, borrow_array_type};
use hugr_core::types::type_param::TypeParam;
use hugr_core::types::{
FuncValueType, PolyFuncTypeRV, Signature, Type, TypeArg, TypeBound, TypeRow,
};
use hugr_core::{Extension, Hugr, HugrView, Node, hugr::IdentList};
use itertools::Itertools;
use rstest::rstest;
use crate::replace_types::handlers::{DISCARD_TO_UNIT_PREFIX, MAKE_NONE_PREFIX, UNWRAP_PREFIX};
use crate::replace_types::{LinearizeError, Linearizer, NodeTemplate, ReplaceTypesError};
use crate::{ComposablePass, ReplaceTypes};
const LIN_T: &str = "Lin";
const COPY_T: &str = "Copy";
struct NWayCopySigFn(Type);
impl CustomSignatureFunc for NWayCopySigFn {
fn compute_signature<'o, 'a: 'o>(
&'a self,
arg_values: &[TypeArg],
_def: &'o OpDef,
) -> Result<PolyFuncTypeRV, SignatureError> {
let [TypeArg::BoundedNat(n)] = arg_values else {
panic!()
};
let outs = vec![self.0.clone(); *n as usize];
Ok(FuncValueType::new([self.0.clone()], outs).into())
}
fn static_params(&self) -> &[TypeParam] {
const JUST_NAT: &[TypeParam] = &[TypeParam::max_nat_type()];
JUST_NAT
}
}
fn ext_lowerer() -> (Arc<Extension>, ReplaceTypes) {
let e = Extension::new_arc(
IdentList::new_unchecked("TestExt"),
Version::new(0, 0, 0),
|e, w| {
let lin = Type::new_extension(
e.add_type(LIN_T.into(), vec![], String::new(), TypeDefBound::any(), w)
.unwrap()
.instantiate([])
.unwrap(),
);
e.add_type(
COPY_T.into(),
vec![],
String::new(),
TypeDefBound::copyable(),
w,
)
.unwrap()
.instantiate([])
.unwrap();
e.add_op(
"discard".into(),
String::new(),
Signature::new([lin.clone()], []),
w,
)
.unwrap();
e.add_op(
"copy".into(),
String::new(),
SignatureFunc::CustomFunc(Box::new(NWayCopySigFn(lin))),
w,
)
.unwrap();
},
);
let lin_custom_t = e.get_type(LIN_T).unwrap().instantiate([]).unwrap();
let copy_op = ExtensionOp::new(e.get_op("copy").unwrap().clone(), [2.into()]).unwrap();
let discard_op = ExtensionOp::new(e.get_op("discard").unwrap().clone(), []).unwrap();
let mut lowerer = ReplaceTypes::default();
let usize_custom_t = usize_t().as_extension().unwrap().clone();
lowerer.set_replace_type(usize_custom_t, Type::new_extension(lin_custom_t.clone()));
lowerer
.linearizer_mut()
.register_simple(
lin_custom_t,
NodeTemplate::SingleOp(copy_op.into()),
NodeTemplate::SingleOp(discard_op.into()),
)
.unwrap();
(e, lowerer)
}
#[test]
fn single_values() {
let (_e, lowerer) = ext_lowerer();
let mut outer = DFGBuilder::new(inout_sig(
vec![usize_t(); 2],
vec![usize_t(), borrow_array_type(2, usize_t())],
))
.unwrap();
let [inp, _] = outer.input_wires_arr();
let new_array = outer
.add_dataflow_op(BArrayOpDef::new_array.to_concrete(usize_t(), 2), [inp, inp])
.unwrap();
let [arr] = new_array.outputs_arr();
let mut h = outer.finish_hugr_with_outputs([inp, arr]).unwrap();
assert!(lowerer.run(&mut h).unwrap());
let ext_ops = h
.entry_descendants()
.filter_map(|n| h.get_optype(n).as_extension_op());
let mut counts = HashMap::<OpName, u32>::new();
for e in ext_ops {
*counts.entry(e.qualified_id()).or_default() += 1;
}
assert_eq!(
counts,
HashMap::from([
("TestExt.copy".into(), 2),
("TestExt.discard".into(), 1),
("collections.borrow_arr.new_array".into(), 1)
])
);
}
fn copy_n_discard_one(ty: Type, n: usize) -> (Hugr, Node) {
let mut outer = DFGBuilder::new(inout_sig([ty.clone()], vec![ty.clone(); n - 1])).unwrap();
let [inp] = outer.input_wires_arr();
let inner = outer
.dfg_builder(inout_sig([ty], []), [inp])
.unwrap()
.finish_with_outputs([])
.unwrap();
let h = outer.finish_hugr_with_outputs(vec![inp; n - 1]).unwrap();
(h, inner.node())
}
#[rstest]
fn sums_2way_copy(#[values(2, 3, 4)] num_copies: usize) {
let (mut h, inner) = copy_n_discard_one(option_type([usize_t()]).into(), num_copies);
let (e, lowerer) = ext_lowerer();
assert!(lowerer.run(&mut h).unwrap());
let lin_t = Type::from(e.get_type(LIN_T).unwrap().instantiate([]).unwrap());
let sum_ty: Type = option_type([lin_t.clone()]).into();
let count_tags = |n| h.children(n).filter(|n| h.get_optype(*n).is_tag()).count();
for (dfg, num_tags, expected_ext_ops) in [
(inner.node(), 0, vec!["TestExt.discard"]),
(
h.entrypoint(),
num_copies,
vec!["TestExt.copy"; num_copies - 1],
), ] {
let [(cond_node, cond)] = h
.children(dfg)
.filter_map(|n| h.get_optype(n).as_conditional().map(|c| (n, c)))
.collect_array()
.unwrap();
assert_eq!(
cond.signature().output(),
&TypeRow::from(vec![sum_ty.clone(); num_tags])
);
let [case0, case1] = h.children(cond_node).collect_array().unwrap();
assert_eq!(h.children(case0).count(), 2 + num_tags); assert_eq!(count_tags(case0), num_tags);
assert_eq!(h.children(case1).count(), 3 + num_tags); assert_eq!(count_tags(case1), num_tags);
let ext_ops = h
.descendants(case1)
.filter_map(|n| {
h.get_optype(n)
.as_extension_op()
.map(ExtensionOp::qualified_id)
})
.collect_vec();
assert_eq!(ext_ops, expected_ext_ops);
}
}
#[rstest]
fn sum_nway_copy(#[values(2, 5, 9)] num_copies: usize) {
let i8_t = || INT_TYPES[3].clone();
let sum_ty = Type::new_sum([vec![i8_t()], vec![usize_t(); 2]]);
let (mut h, inner) = copy_n_discard_one(sum_ty, num_copies);
let (e, _) = ext_lowerer();
let mut lowerer = ReplaceTypes::default();
let lin_t_def = e.get_type(LIN_T).unwrap();
lowerer.set_replace_type(
usize_t().as_extension().unwrap().clone(),
lin_t_def.instantiate([]).unwrap().into(),
);
let opdef = e.get_op("copy").unwrap();
let opdef2 = opdef.clone();
lowerer
.linearizer_mut()
.register_callback(lin_t_def, move |args, num_outs, _| {
assert!(args.is_empty());
Ok(NodeTemplate::SingleOp(
ExtensionOp::new(opdef2.clone(), [(num_outs as u64).into()])
.unwrap()
.into(),
))
});
assert!(lowerer.run(&mut h).unwrap());
let lin_t = Type::from(e.get_type(LIN_T).unwrap().instantiate([]).unwrap());
let sum_ty = Type::new_sum([vec![i8_t()], vec![lin_t.clone(); 2]]);
let count_tags = |n| h.children(n).filter(|n| h.get_optype(*n).is_tag()).count();
for (dfg, num_tags) in [(inner.node(), 0), (h.entrypoint(), num_copies)] {
let [cond] = h
.children(dfg)
.filter(|n| h.get_optype(*n).is_conditional())
.collect_array()
.unwrap();
let [case0, case1] = h.children(cond).collect_array().unwrap();
let out_row = vec![sum_ty.clone(); num_tags].into();
assert_eq!(h.children(case0).count(), 2 + num_tags); assert_eq!(count_tags(case0), num_tags);
let case0 = h.get_optype(case0).as_case().unwrap();
assert_eq!(case0.signature.io(), (&vec![i8_t()].into(), &out_row));
assert_eq!(h.children(case1).count(), 4 + num_tags); assert_eq!(count_tags(case1), num_tags);
let ext_ops = h
.children(case1)
.filter_map(|n| h.get_optype(n).as_extension_op())
.collect_vec();
let expected_op = ExtensionOp::new(opdef.clone(), [(num_tags as u64).into()]).unwrap();
assert_eq!(ext_ops, vec![&expected_op; 2]);
let case1 = h.get_optype(case1).as_case().unwrap();
assert_eq!(
case1.signature.io(),
(&vec![lin_t.clone(); 2].into(), &out_row)
);
}
}
#[test]
fn bad_sig() {
let (ext, _) = ext_lowerer();
let lin_ct = ext.get_type(LIN_T).unwrap().instantiate([]).unwrap();
let lin_t = Type::from(lin_ct.clone());
let copy3 = OpType::from(
ExtensionOp::new(ext.get_op("copy").unwrap().clone(), [3.into()]).unwrap(),
);
let copy2 = ExtensionOp::new(ext.get_op("copy").unwrap().clone(), [2.into()]).unwrap();
let discard = ExtensionOp::new(ext.get_op("discard").unwrap().clone(), []).unwrap();
let mut replacer = ReplaceTypes::default();
replacer.set_replace_type(usize_t().as_extension().unwrap().clone(), lin_t.clone());
let bad_copy = replacer.linearizer_mut().register_simple(
lin_ct.clone(),
NodeTemplate::SingleOp(copy3.clone()),
NodeTemplate::SingleOp(discard.clone().into()),
);
let sig3 = Some(Signature::new([lin_t.clone()], vec![lin_t.clone(); 3]));
assert_eq!(
bad_copy,
Err(LinearizeError::WrongSignature {
typ: Box::new(lin_t.clone()),
num_outports: 2,
sig: sig3.clone().map(Box::new)
})
);
let bad_discard = replacer.linearizer_mut().register_simple(
lin_ct.clone(),
NodeTemplate::SingleOp(copy2.into()),
NodeTemplate::SingleOp(copy3.clone()),
);
assert_eq!(
bad_discard,
Err(LinearizeError::WrongSignature {
typ: Box::new(lin_t.clone()),
num_outports: 0,
sig: sig3.clone().map(Box::new)
})
);
replacer
.linearizer_mut()
.register_callback(ext.get_type(LIN_T).unwrap(), move |_args, _, _| {
Ok(NodeTemplate::SingleOp(copy3.clone()))
});
let dfb = DFGBuilder::new(inout_sig([usize_t()], vec![usize_t(); 2])).unwrap();
let [inp] = dfb.input_wires_arr();
let mut h = dfb.finish_hugr_with_outputs([inp, inp]).unwrap();
assert_eq!(
replacer.run(&mut h),
Err(ReplaceTypesError::LinearizeError(
LinearizeError::WrongSignature {
typ: Box::new(lin_t.clone()),
num_outports: 2,
sig: sig3.clone().map(Box::new)
}
))
);
}
#[rstest]
fn call_in_array(#[values(true, false)] use_linking: bool) {
let (e, _) = ext_lowerer();
let lin_ct = e.get_type(LIN_T).unwrap().instantiate([]).unwrap();
let lin_t: Type = lin_ct.clone().into();
let mut dfb = DFGBuilder::new(inout_sig([usize_t()], [])).unwrap();
let discard_fn = {
let mut mb = dfb.module_root_builder();
let mut fb = mb
.define_function_vis(
"drop",
Signature::new([lin_t.clone()], []),
Visibility::Public,
)
.unwrap();
let ins = fb.input_wires();
fb.add_dataflow_op(
ExtensionOp::new(e.get_op("discard").unwrap().clone(), []).unwrap(),
ins,
)
.unwrap();
fb.finish_with_outputs([]).unwrap()
}
.node();
let backup = dfb.finish_hugr().unwrap();
let mut lower_discard_to_call = ReplaceTypes::default();
if use_linking {
lower_discard_to_call
.linearizer_mut()
.register_simple(
lin_ct.clone(),
NodeTemplate::CompoundOp(Box::new({
std::mem::take(
DFGBuilder::new(inout_sig([lin_t.clone()], vec![lin_t.clone(); 2]))
.unwrap()
.hugr_mut(),
)
})),
NodeTemplate::linked_hugr({
let mut dfb = DFGBuilder::new(inout_sig([lin_t.clone()], [])).unwrap();
let drop_fn = dfb
.module_root_builder()
.declare("drop", inout_sig([lin_t.clone()], []).into())
.unwrap();
let ins = dfb.input_wires();
let call = dfb.call(&drop_fn, &[], ins).unwrap();
dfb.finish_hugr_with_outputs(call.outputs()).unwrap()
}),
)
.unwrap();
} else {
#[expect(deprecated)] lower_discard_to_call
.linearizer_mut()
.register_simple(
lin_ct.clone(),
NodeTemplate::Call(backup.entrypoint(), vec![]), NodeTemplate::Call(discard_fn, vec![]),
)
.unwrap();
};
{
let mut lowerer = lower_discard_to_call.clone();
lowerer.set_replace_type(usize_t().as_extension().unwrap().clone(), lin_t.clone());
let mut h = backup.clone();
lowerer.run(&mut h).unwrap();
assert_eq!(h.output_neighbours(discard_fn).count(), 1);
}
lower_discard_to_call.set_replace_type(
usize_t().as_extension().unwrap().clone(),
array_type(4, lin_ct.into()),
);
let mut h = backup.clone();
let r = lower_discard_to_call.run(&mut h);
if use_linking {
r.unwrap();
h.validate().unwrap();
} else {
r.unwrap();
h.validate().unwrap();
let disc = h
.children(h.module_root())
.find(|n| {
h.get_optype(*n)
.as_func_defn()
.is_some_and(|fd| fd.func_name().contains(DISCARD_TO_UNIT_PREFIX))
})
.unwrap();
let call = h
.descendants(disc)
.filter(|n| h.get_optype(*n).is_call())
.exactly_one()
.ok()
.unwrap();
assert_eq!(h.static_source(call), Some(disc)); }
}
#[test]
fn use_in_op_callback() {
let (e, mut lowerer) = ext_lowerer();
let drop_ext = Extension::new_arc(
IdentList::new_unchecked("DropExt"),
Version::new(0, 0, 0),
|e, w| {
e.add_op(
"drop".into(),
String::new(),
PolyFuncTypeRV::new(
[TypeBound::Linear.into()], Signature::new([Type::new_var_use(0, TypeBound::Linear)], vec![]),
),
w,
)
.unwrap();
},
);
let drop_op = drop_ext.get_op("drop").unwrap();
lowerer.set_replace_parametrized_op(drop_op, |args, rt| {
let [TypeArg::Runtime(ty)] = args else {
panic!("Expected just one type")
};
Ok(Some(rt.get_linearizer().copy_discard_op(ty, 0)?))
});
let build_hugr = |ty: Type| {
let mut dfb = DFGBuilder::new(Signature::new([ty.clone()], [])).unwrap();
let [inp] = dfb.input_wires_arr();
let drop_op = drop_ext
.instantiate_extension_op("drop", [ty.into()])
.unwrap();
dfb.add_dataflow_op(drop_op, [inp]).unwrap();
dfb.finish_hugr().unwrap()
};
let lin_t = Type::from(e.get_type(LIN_T).unwrap().instantiate([]).unwrap());
let mut h = build_hugr(Type::new_tuple(vec![lin_t.clone(); 2]));
lowerer.run(&mut h).unwrap();
h.validate().unwrap();
let mut exts = h.nodes().filter_map(|n| h.get_optype(n).as_extension_op());
assert_eq!(exts.clone().count(), 2);
assert!(exts.all(|eo| eo.qualified_id() == "TestExt.discard"));
let mut h = build_hugr(borrow_array_type(4, lin_t));
lowerer.run(&mut h).unwrap();
h.validate().unwrap();
let mut exts = h.nodes().filter_map(|n| h.get_optype(n).as_extension_op());
assert!(exts.any(|eo| eo.qualified_id() == "collections.borrow_arr.discard_all_borrowed"));
let mut h = build_hugr(borrow_array_type(4, usize_t()));
lowerer.run(&mut h).unwrap();
h.validate().unwrap();
let mut exts = h.nodes().filter_map(|n| h.get_optype(n).as_extension_op());
assert!(exts.any(|eo| eo.qualified_id() == "collections.borrow_arr.discard_all_borrowed"));
let mut h = build_hugr(qb_t());
assert_eq!(
lowerer.run(&mut h).unwrap_err(),
ReplaceTypesError::LinearizeError(LinearizeError::NeedCopyDiscard(Box::new(qb_t())))
);
let mut h = build_hugr(borrow_array_type(4, qb_t()));
assert_eq!(
lowerer.run(&mut h).unwrap_err(),
ReplaceTypesError::LinearizeError(LinearizeError::NeedCopyDiscard(Box::new(qb_t())))
);
}
#[rstest]
#[case([borrow_array_type(2, usize_t())])]
#[case([borrow_array_type(2, usize_t()), borrow_array_type(4, usize_t())])]
fn test_copy_borrow_array<const N: usize>(#[case] tys: [Type; N]) {
let (inp, out, mut h) = {
let mut dfb = DFGBuilder::new(Signature::new(
Vec::from_iter(tys.clone()),
tys.clone().into_iter().chain(tys.clone()).collect_vec(),
))
.unwrap();
(dfb.input(), dfb.output(), std::mem::take(dfb.hugr_mut()))
};
for (n, _) in tys.iter().enumerate() {
h.connect(inp.node(), n, out.node(), n);
h.connect(inp.node(), n, out.node(), n + tys.len());
}
assert!(matches!(
h.validate(),
Err(ValidationError::TooManyConnections { .. })
));
let (_e, lowerer) = ext_lowerer();
lowerer.run(&mut h).unwrap();
h.validate().unwrap();
for prefix in [UNWRAP_PREFIX, MAKE_NONE_PREFIX] {
assert_eq!(
h.children(h.module_root())
.filter(|n| match h.get_optype(*n) {
OpType::FuncDecl(_) => panic!("Unexpected FuncDecl"),
OpType::FuncDefn(fd) => fd.func_name().contains(prefix),
_ => false,
})
.count(),
1,
"Found multiple {prefix} funcs"
);
}
}
}