use std::collections::HashMap;
use crate::measurements::{
BasisRotationWrapper, CheatedBasisRotationWrapper, CheatedWrapper, ClassicalRegisterWrapper,
};
use crate::{QoqoError, QOQO_VERSION};
use bincode::{deserialize, serialize};
use pyo3::exceptions::{PyRuntimeError, PyTypeError, PyValueError};
use pyo3::prelude::*;
use pyo3::types::PyByteArray;
use pyo3::types::PyType;
use pyo3::PyObjectProtocol;
use roqoqo::measurements;
use roqoqo::measurements::Measure;
use roqoqo::QuantumProgram;
use roqoqo::ROQOQO_VERSION;
#[pyclass(name = "QuantumProgram", module = "qoqo")]
#[derive(Clone, Debug, PartialEq, serde::Serialize, serde::Deserialize)]
pub struct QuantumProgramWrapper {
pub internal: QuantumProgram,
}
#[pymethods]
impl QuantumProgramWrapper {
#[new]
pub fn new(measurement: &PyAny, input_parameter_names: Vec<String>) -> PyResult<Self> {
if let Ok(try_downcast) = measurement.extract::<BasisRotationWrapper>() {
return Ok(Self {
internal: QuantumProgram::BasisRotation {
measurement: try_downcast.internal,
input_parameter_names,
},
});
}
if let Ok(try_downcast) = measurement.extract::<CheatedBasisRotationWrapper>() {
return Ok(Self {
internal: QuantumProgram::CheatedBasisRotation {
measurement: try_downcast.internal,
input_parameter_names,
},
});
}
if let Ok(try_downcast) = measurement.extract::<CheatedWrapper>() {
return Ok(Self {
internal: QuantumProgram::Cheated {
measurement: try_downcast.internal,
input_parameter_names,
},
});
}
if let Ok(try_downcast) = measurement.extract::<ClassicalRegisterWrapper>() {
return Ok(Self {
internal: QuantumProgram::ClassicalRegister {
measurement: try_downcast.internal,
input_parameter_names,
},
});
}
let get_measurement_internal = measurement
.call_method0("_internal_to_bincode")
.map_err(|_| PyTypeError::new_err("measurement is not of type Measurement. Are you using different versions of roqoqo?"))?;
let (name, encoded) = get_measurement_internal
.extract::<(&str, &[u8])>()
.map_err(|_| PyTypeError::new_err("measurement is not of type Measurement. Are you using different versions of roqoqo?"))?;
match name {
"BasisRotation" => {
let measure: measurements::BasisRotation = deserialize(encoded).map_err(|_| PyTypeError::new_err("measurement is not of type Measurement. Are you using different versions of roqoqo?"))?;
Ok( Self{internal: QuantumProgram::BasisRotation{measurement: measure, input_parameter_names}})
},
"CheatedBasisRotation" => {
let measure: measurements::CheatedBasisRotation = deserialize(encoded).map_err(|_| PyTypeError::new_err("measurement is not of type Measurement. Are you using different versions of roqoqo?"))?;
Ok( Self{internal: QuantumProgram::CheatedBasisRotation{measurement: measure, input_parameter_names}})
},
"Cheated" => {
let measure: measurements::Cheated = deserialize(encoded).map_err(|_| PyTypeError::new_err("measurement is not of type Measurement. Are you using different versions of roqoqo?"))?;
Ok( Self{internal: QuantumProgram::Cheated{measurement: measure, input_parameter_names}})
}
"ClassicalRegister" => {
let measure: measurements::ClassicalRegister = deserialize(encoded).map_err(|_| PyTypeError::new_err("measurement is not of type Measurement. Are you using different versions of roqoqo?"))?;
Ok( Self{internal: QuantumProgram::ClassicalRegister{measurement: measure, input_parameter_names}})
}
_ => Err(PyTypeError::new_err("measurement is not of type Measurement. Are you using different versions of roqoqo?"))
}
}
pub fn measurement(&self) -> PyObject {
match self.internal.clone() {
QuantumProgram::BasisRotation {
measurement,
input_parameter_names: _,
} => Python::with_gil(|py| -> PyObject {
let pyref: Py<BasisRotationWrapper> = Py::new(
py,
BasisRotationWrapper {
internal: measurement.clone(),
},
)
.unwrap();
pyref.to_object(py)
}),
QuantumProgram::CheatedBasisRotation {
measurement,
input_parameter_names: _,
} => Python::with_gil(|py| -> PyObject {
let pyref: Py<CheatedBasisRotationWrapper> = Py::new(
py,
CheatedBasisRotationWrapper {
internal: measurement.clone(),
},
)
.unwrap();
pyref.to_object(py)
}),
QuantumProgram::Cheated {
measurement,
input_parameter_names: _,
} => Python::with_gil(|py| -> PyObject {
let pyref: Py<CheatedWrapper> = Py::new(
py,
CheatedWrapper {
internal: measurement.clone(),
},
)
.unwrap();
pyref.to_object(py)
}),
QuantumProgram::ClassicalRegister {
measurement,
input_parameter_names: _,
} => Python::with_gil(|py| -> PyObject {
let pyref: Py<ClassicalRegisterWrapper> = Py::new(
py,
ClassicalRegisterWrapper {
internal: measurement.clone(),
},
)
.unwrap();
pyref.to_object(py)
}),
}
}
pub fn input_parameter_names(&self) -> Vec<String> {
match self.internal.clone() {
QuantumProgram::BasisRotation {
measurement: _,
input_parameter_names,
} => input_parameter_names,
QuantumProgram::CheatedBasisRotation {
measurement: _,
input_parameter_names,
} => input_parameter_names,
QuantumProgram::Cheated {
measurement: _,
input_parameter_names,
} => input_parameter_names,
QuantumProgram::ClassicalRegister {
measurement: _,
input_parameter_names,
} => input_parameter_names,
}
}
pub fn run(&self, backend: Py<PyAny>, parameters: Option<Vec<f64>>) -> PyResult<Py<PyAny>> {
let parameters = parameters.unwrap_or_default();
match &self.internal{
QuantumProgram::BasisRotation{measurement, input_parameter_names } => {
if parameters.len() != input_parameter_names.len() { return Err(PyValueError::new_err( format!("Wrong number of parameters {} parameters expected {} parameters given", input_parameter_names.len(), parameters.len())))};
let substituted_parameters: HashMap<String, f64> = input_parameter_names.iter().zip(parameters.iter()).map(|(key, value)| (key.clone(), *value)).collect();
let substituted_measurement = measurement.substitute_parameters(
substituted_parameters
).map_err(|err| PyRuntimeError::new_err(format!("Applying parameters failed {:?}", err)))?;
Python::with_gil(|py| -> PyResult<Py<PyAny>> {
backend.call_method1(py, "run_measurement", (BasisRotationWrapper{internal: substituted_measurement}, ))
}) }
QuantumProgram::CheatedBasisRotation{measurement, input_parameter_names } => {
if parameters.len() != input_parameter_names.len() { return Err(PyValueError::new_err( format!("Wrong number of parameters {} parameters expected {} parameters given", input_parameter_names.len(), parameters.len())))};
let substituted_parameters: HashMap<String, f64> = input_parameter_names.iter().zip(parameters.iter()).map(|(key, value)| (key.clone(), *value)).collect();
let substituted_measurement = measurement.substitute_parameters(
substituted_parameters
).map_err(|err| PyRuntimeError::new_err(format!("Applying parameters failed {:?}", err)))?;
Python::with_gil(|py| -> PyResult<Py<PyAny>> {
backend.call_method1(py, "run_measurement", (CheatedBasisRotationWrapper{internal: substituted_measurement}, ))
})
}
QuantumProgram::Cheated{measurement, input_parameter_names } => {
if parameters.len() != input_parameter_names.len() { return Err(PyValueError::new_err( format!("Wrong number of parameters {} parameters expected {} parameters given", input_parameter_names.len(), parameters.len())))};
let substituted_parameters: HashMap<String, f64> = input_parameter_names.iter().zip(parameters.iter()).map(|(key, value)| (key.clone(), *value)).collect();
let substituted_measurement = measurement.substitute_parameters(
substituted_parameters
).map_err(|err| PyRuntimeError::new_err(format!("Applying parameters failed {:?}", err)))?;
Python::with_gil(|py| -> PyResult<Py<PyAny>> {
backend.call_method1(py, "run_measurement", (CheatedWrapper{internal: substituted_measurement}, ))
}) }
_ => Err(PyTypeError::new_err("A quantum programm returning classical registeres cannot be executed by `run` use `run_registers` instead".to_string()))
}
}
pub fn run_registers(
&self,
backend: Py<PyAny>,
parameters: Option<Vec<f64>>,
) -> PyResult<Py<PyAny>> {
let parameters = parameters.unwrap_or_default();
match &self.internal{
QuantumProgram::ClassicalRegister{measurement, input_parameter_names } => {
if parameters.len() != input_parameter_names.len() { return Err(PyValueError::new_err( format!("Wrong number of parameters {} parameters expected {} parameters given", input_parameter_names.len(), parameters.len())))};
let substituted_parameters: HashMap<String, f64> = input_parameter_names.iter().zip(parameters.iter()).map(|(key, value)| (key.clone(), *value)).collect();
let substituted_measurement = measurement.substitute_parameters(
substituted_parameters
).map_err(|err| PyRuntimeError::new_err(format!("Applying parameters failed {:?}", err)))?;
Python::with_gil(|py| -> PyResult<Py<PyAny>> {
backend.call_method1(py, "run_measurement_registers", (ClassicalRegisterWrapper{internal: substituted_measurement}, ))
}) },
_ => Err(PyTypeError::new_err("A quantum programm returning expectation values cannot be executed by `run_registers` use `run` instead".to_string()))
}
}
pub fn __copy__(&self) -> QuantumProgramWrapper {
self.clone()
}
pub fn __deepcopy__(&self, _memodict: Py<PyAny>) -> QuantumProgramWrapper {
self.clone()
}
fn _qoqo_versions(&self) -> (String, String) {
let mut rsplit = ROQOQO_VERSION.split('.').take(2);
let mut qsplit = QOQO_VERSION.split('.').take(2);
let rver = format!(
"{}.{}",
rsplit.next().expect("ROQOQO_VERSION badly formatted"),
rsplit.next().expect("ROQOQO_VERSION badly formatted")
);
let qver = format!(
"{}.{}",
qsplit.next().expect("QOQO_VERSION badly formatted"),
qsplit.next().expect("QOQO_VERSION badly formatted")
);
(rver, qver)
}
pub fn to_bincode(&self) -> PyResult<Py<PyByteArray>> {
let serialized = serialize(&self.internal)
.map_err(|_| PyValueError::new_err("Cannot serialize QuantumProgram to bytes"))?;
let b: Py<PyByteArray> = Python::with_gil(|py| -> Py<PyByteArray> {
PyByteArray::new(py, &serialized[..]).into()
});
Ok(b)
}
#[allow(unused_variables)]
#[classmethod]
pub fn from_bincode(cls: &PyType, input: &PyAny) -> PyResult<Self> {
let bytes = input
.extract::<Vec<u8>>()
.map_err(|_| PyTypeError::new_err("Input cannot be converted to byte array"))?;
Ok(Self {
internal: deserialize(&bytes[..]).map_err(|_| {
PyValueError::new_err("Input cannot be deserialized to QuantumProgram")
})?,
})
}
fn to_json(&self) -> PyResult<String> {
let serialized = serde_json::to_string(&self.internal)
.map_err(|_| PyValueError::new_err("Cannot serialize QuantumProgram to json"))?;
Ok(serialized)
}
#[allow(unused_variables)]
#[classmethod]
fn from_json(cls: &PyType, input: &str) -> PyResult<Self> {
Ok(Self {
internal: serde_json::from_str(input).map_err(|_| {
PyValueError::new_err("Input cannot be deserialized to QuantumProgram")
})?,
})
}
}
#[pyproto]
impl PyObjectProtocol for QuantumProgramWrapper {
fn __richcmp__(&self, other: Py<PyAny>, op: pyo3::class::basic::CompareOp) -> PyResult<bool> {
let other = Python::with_gil(|py| -> Result<QuantumProgram, QoqoError> {
let other_ref = other.as_ref(py);
convert_into_quantum_program(other_ref)
});
match op {
pyo3::class::basic::CompareOp::Eq => match other {
Ok(qp) => Ok(self.internal == qp),
_ => Ok(false),
},
pyo3::class::basic::CompareOp::Ne => match other {
Ok(qp) => Ok(self.internal != qp),
_ => Ok(true),
},
_ => Err(pyo3::exceptions::PyNotImplementedError::new_err(
"Other comparison not implemented",
)),
}
}
}
pub fn convert_into_quantum_program(input: &PyAny) -> Result<QuantumProgram, QoqoError> {
if let Ok(try_downcast) = input.extract::<QuantumProgramWrapper>() {
return Ok(try_downcast.internal);
}
let get_version = input
.call_method0("_qoqo_versions")
.map_err(|_| QoqoError::CannotExtractObject)?;
let version = get_version
.extract::<(&str, &str)>()
.map_err(|_| QoqoError::CannotExtractObject)?;
let mut rsplit = ROQOQO_VERSION.split('.').take(2);
let mut qsplit = QOQO_VERSION.split('.').take(2);
let rver = format!(
"{}.{}",
rsplit.next().expect("ROQOQO_VERSION badly formatted"),
rsplit.next().expect("ROQOQO_VERSION badly formatted")
);
let qver = format!(
"{}.{}",
qsplit.next().expect("QOQO_VERSION badly formatted"),
qsplit.next().expect("QOQO_VERSION badly formatted")
);
let test_version: (&str, &str) = (rver.as_str(), qver.as_str());
if version == test_version {
let get_bytes = input
.call_method0("to_bincode")
.map_err(|_| QoqoError::CannotExtractObject)?;
let bytes = get_bytes
.extract::<Vec<u8>>()
.map_err(|_| QoqoError::CannotExtractObject)?;
deserialize(&bytes[..]).map_err(|_| QoqoError::CannotExtractObject)
} else {
Err(QoqoError::VersionMismatch)
}
}