Skip to main content

tket/extension/
modifier.rs

1//! This module defines a Hugr extension for modifier operations.
2//! These operations modify circuits by applying modifiers: control, dagger, or power.
3use 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/// Types of modifers.
20#[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    /// Control modifier.
37    ControlModifier,
38    /// Dagger modifier.
39    DaggerModifier,
40    /// Power modifier.
41    PowerModifier,
42}
43
44/// Identifier for the `tket.modifier` extension.
45pub const MODIFIER_EXTENSION_ID: ExtensionId = ExtensionId::new_unchecked("tket.modifier");
46/// Version of the `tket.modifier` extension.
47pub const MODIFIER_VERSION: Version = Version::new(0, 1, 0);
48
49lazy_static! {
50    /// The extension definition for modifier operations.
51    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
77/// Identifier for the `ControlModifier` operation.
78pub const CONTROL_OP_ID: OpName = OpName::new_inline("ControlModifier");
79/// Identifier for the `DaggerModifier` operation.
80pub const DAGGER_OP_ID: OpName = OpName::new_inline("DaggerModifier");
81/// Identifier for the `PowerModifier` operation.
82pub 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    // [TODO]: Do we need this?
129    // fn post_opdef(&self, _def: &mut OpDef);
130}
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            // vec![loaded_func]
239        } 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}