1use lazy_static::lazy_static;
4use std::sync::{Arc, Weak};
5
6use hugr::{
7 Extension,
8 extension::{
9 ExtensionId, OpDef, SignatureFunc, Version,
10 simple_op::{MakeOpDef, OpLoadError},
11 },
12 ops::OpName,
13};
14use serde::{Deserialize, Serialize};
15use strum::{EnumIter, EnumString, IntoStaticStr};
16
17use crate::modifier::{control::ModifierControl, dagger::ModifierDagger, power::ModifierPower};
18
19#[derive(
21 Clone,
22 Copy,
23 Debug,
24 Serialize,
25 Deserialize,
26 Hash,
27 PartialEq,
28 Eq,
29 PartialOrd,
30 Ord,
31 EnumIter,
32 IntoStaticStr,
33 EnumString,
34)]
35pub enum Modifier {
36 ControlModifier,
38 DaggerModifier,
40 PowerModifier,
42}
43
44pub const MODIFIER_EXTENSION_ID: ExtensionId = ExtensionId::new_unchecked("tket.modifier");
46pub const MODIFIER_VERSION: Version = Version::new(0, 1, 0);
48
49lazy_static! {
50 pub static ref MODIFIER_EXTENSION: Arc<Extension> = {
52 Extension::new_arc(MODIFIER_EXTENSION_ID, MODIFIER_VERSION, |modifier, extension_ref| {
53 modifier.add_op(
54 CONTROL_OP_ID,
55 "Quantum control operation".to_string(),
56 ModifierControl::signature(),
57 extension_ref,
58 ).unwrap();
59
60 modifier.add_op(
61 DAGGER_OP_ID,
62 "Dagger Operator".to_string(),
63 ModifierDagger::signature(),
64 extension_ref,
65 ).unwrap();
66
67 modifier.add_op(
68 POWER_OP_ID,
69 "Power Operator".to_string(),
70 ModifierPower::signature(),
71 extension_ref,
72 ).unwrap();
73 }
74 )};
75}
76
77pub const CONTROL_OP_ID: OpName = OpName::new_inline("ControlModifier");
79pub const DAGGER_OP_ID: OpName = OpName::new_inline("DaggerModifier");
81pub const POWER_OP_ID: OpName = OpName::new_inline("PowerModifier");
83
84impl MakeOpDef for Modifier {
85 fn opdef_id(&self) -> OpName {
86 match self {
87 Modifier::ControlModifier => CONTROL_OP_ID.clone(),
88 Modifier::DaggerModifier => DAGGER_OP_ID.clone(),
89 Modifier::PowerModifier => POWER_OP_ID.clone(),
90 }
91 }
92
93 fn from_def(op_def: &OpDef) -> Result<Self, OpLoadError>
94 where
95 Self: Sized,
96 {
97 hugr::extension::simple_op::try_from_name(op_def.name(), op_def.extension_id())
98 }
99
100 fn init_signature(&self, _extension_ref: &std::sync::Weak<hugr::Extension>) -> SignatureFunc {
101 match self {
102 Modifier::ControlModifier => ModifierControl::signature(),
103 Modifier::DaggerModifier => ModifierDagger::signature(),
104 Modifier::PowerModifier => ModifierPower::signature(),
105 }
106 }
107
108 fn extension_ref(&self) -> Weak<hugr::Extension> {
109 Arc::downgrade(&MODIFIER_EXTENSION)
110 }
111
112 fn extension(&self) -> ExtensionId {
113 MODIFIER_EXTENSION_ID.to_owned()
114 }
115
116 fn description(&self) -> String {
117 match self {
118 Modifier::ControlModifier => {
119 "Generates a quantum-controlled circuit from a circuit.".into()
120 }
121 Modifier::DaggerModifier => "Dagger operation on a circuit.".into(),
122 Modifier::PowerModifier => {
123 "Generates a circuit that applies a circuit many times.".into()
124 }
125 }
126 }
127
128 }
131
132#[cfg(test)]
133mod test {
134 use super::{
135 CONTROL_OP_ID, DAGGER_OP_ID, MODIFIER_EXTENSION, MODIFIER_EXTENSION_ID, Modifier,
136 POWER_OP_ID,
137 };
138 use cool_asserts::assert_matches;
139 use hugr::{
140 builder::{Dataflow, DataflowSubContainer, HugrBuilder, ModuleBuilder},
141 extension::{
142 OpDef,
143 prelude::{bool_t, qb_t},
144 simple_op::{MakeExtensionOp, MakeOpDef},
145 },
146 ops::{CallIndirect, ExtensionOp},
147 std_extensions::{
148 arithmetic::int_types::{ConstInt, int_type},
149 collections::array::array_type,
150 },
151 types::{Signature, Term, Type},
152 };
153 use rstest::rstest;
154 use std::sync::Arc;
155 use strum::IntoEnumIterator;
156
157 fn get_modifier_opdef(op: Modifier) -> Option<&'static Arc<OpDef>> {
158 MODIFIER_EXTENSION.get_op(&op.op_id())
159 }
160
161 #[test]
162 fn create_modifier_extension() {
163 assert_eq!(MODIFIER_EXTENSION.name(), &MODIFIER_EXTENSION_ID);
164
165 for o in Modifier::iter() {
166 assert_eq!(Modifier::from_def(get_modifier_opdef(o).unwrap()), Ok(o));
167 }
168 }
169
170 fn control_op(inout: Type, other_inputs: Type) -> (ExtensionOp, Signature) {
171 let modified_sig = Signature::new(
172 vec![array_type(1, qb_t()), inout.clone(), other_inputs.clone()],
173 vec![array_type(1, qb_t()), inout.clone()],
174 );
175 let control_op = MODIFIER_EXTENSION
176 .instantiate_extension_op(
177 &CONTROL_OP_ID,
178 [
179 Term::BoundedNat(1),
180 Term::new_list([inout]),
181 Term::new_list([other_inputs]),
182 ],
183 )
184 .unwrap();
185 (control_op, modified_sig)
186 }
187
188 fn dagger_op(inout: Type, other_inputs: Type) -> (ExtensionOp, Signature) {
189 let modified_sig = Signature::new(
190 vec![inout.clone(), other_inputs.clone()],
191 vec![inout.clone()],
192 );
193 let dagger_op = MODIFIER_EXTENSION
194 .instantiate_extension_op(
195 &DAGGER_OP_ID,
196 [Term::new_list([inout]), Term::new_list([other_inputs])],
197 )
198 .unwrap();
199 (dagger_op, modified_sig)
200 }
201
202 fn power_op(inout: Type, other_inputs: Type) -> (ExtensionOp, Signature) {
203 let modified_sig = Signature::new(
204 vec![inout.clone(), other_inputs.clone()],
205 vec![inout.clone()],
206 );
207 let power_op = MODIFIER_EXTENSION
208 .instantiate_extension_op(
209 &POWER_OP_ID,
210 [Term::new_list([inout]), Term::new_list([other_inputs])],
211 )
212 .unwrap();
213 (power_op, modified_sig)
214 }
215
216 #[rstest]
217 #[case(control_op, false)]
218 #[case(dagger_op, false)]
219 #[case(power_op, true)]
220 fn modifier_op(
221 #[case] op_fn: fn(Type, Type) -> (ExtensionOp, Signature),
222 #[case] needs_extra_param: bool,
223 ) {
224 let original_sig = Signature::new([int_type(6), bool_t()], [int_type(6)]);
225 let (control_op, modified_sig) = op_fn(int_type(6), bool_t());
226 let main_sig = modified_sig.clone();
227
228 let mut module = ModuleBuilder::new();
229
230 let decl = module.declare("dummy_decl", original_sig.into()).unwrap();
231
232 let mut main = module.define_function("_main", main_sig).unwrap();
233 let inputs = main.input_wires();
234 let loaded_func = main.load_func(&decl, &[]).unwrap();
235 let modifier_arg = if needs_extra_param {
236 let int = main.add_load_value(ConstInt::new_u(6, 3).unwrap());
237 vec![loaded_func, int]
238 } else {
240 vec![loaded_func]
241 };
242 let modified = main
243 .add_dataflow_op(control_op, modifier_arg)
244 .unwrap()
245 .out_wire(0);
246 let outputs = main
247 .add_dataflow_op(
248 CallIndirect {
249 signature: modified_sig,
250 },
251 [modified].into_iter().chain(inputs),
252 )
253 .unwrap()
254 .outputs();
255
256 main.finish_with_outputs(outputs).unwrap();
257
258 assert_matches!(module.finish_hugr(), Ok(_));
259 }
260}