use crate::CircuitWrapper;
use bincode::{deserialize, serialize};
use pyo3::exceptions::{PyRuntimeError, PyTypeError, PyValueError};
use pyo3::prelude::*;
use pyo3::types::PyByteArray;
use roqoqo::measurements::ClassicalRegister;
use roqoqo::prelude::*;
use roqoqo::Circuit;
use std::collections::HashMap;
#[pyclass(name = "ClassicalRegister", module = "qoqo.measurements")]
#[derive(Clone, Debug)]
pub struct ClassicalRegisterWrapper {
pub internal: ClassicalRegister,
}
#[pymethods]
impl ClassicalRegisterWrapper {
#[new]
#[pyo3(signature=(constant_circuit, circuits))]
pub fn new(constant_circuit: Option<Py<PyAny>>, circuits: Vec<Py<PyAny>>) -> PyResult<Self> {
let mut new_circuits: Vec<Circuit> = Vec::new();
for c in circuits.into_iter() {
let tmp_c = CircuitWrapper::from_pyany(c).map_err(|err| {
PyTypeError::new_err(format!(
"`circuits` argument is not a list of qoqo Circuits: {}",
err
))
})?;
new_circuits.push(tmp_c)
}
let new_constant: Option<Circuit> = match constant_circuit {
None => None,
Some(c) => {
let tmp_c = CircuitWrapper::from_pyany(c).map_err(|err| {
PyTypeError::new_err(format!(
"`constant_circuit` argument is not None or a qoqo Circuit: {}",
err
))
})?;
Some(tmp_c)
}
};
Ok(Self {
internal: ClassicalRegister {
constant_circuit: new_constant,
circuits: new_circuits,
},
})
}
pub fn circuits(&self) -> Vec<CircuitWrapper> {
self.internal
.circuits()
.map(|c| CircuitWrapper {
internal: c.clone(),
})
.collect()
}
pub fn constant_circuit(&self) -> Option<CircuitWrapper> {
self.internal
.constant_circuit()
.clone()
.map(|c| CircuitWrapper { internal: c })
}
pub fn measurement_type(&self) -> &'static str {
"ClassicalRegister"
}
pub fn substitute_parameters(
&self,
substituted_parameters: HashMap<String, f64>,
) -> PyResult<Self> {
Ok(Self {
internal: self
.internal
.substitute_parameters(substituted_parameters)
.map_err(|x| {
PyRuntimeError::new_err(format!(
"Error substituting symbolic parameters {:?}",
x
))
})?,
})
}
pub fn _internal_to_bincode(&self) -> PyResult<(&'static str, Py<PyByteArray>)> {
let serialized = serialize(&self.internal)
.map_err(|_| PyValueError::new_err("Cannot serialize ClassicalRegister to bytes"))?;
let b: Py<PyByteArray> = Python::with_gil(|py| -> Py<PyByteArray> {
PyByteArray::new(py, &serialized[..]).into()
});
Ok(("ClassicalRegister", b))
}
pub fn to_bincode(&self) -> PyResult<Py<PyByteArray>> {
let serialized = serialize(&self.internal)
.map_err(|_| PyValueError::new_err("Cannot serialize ClassicalRegister to bytes"))?;
let b: Py<PyByteArray> = Python::with_gil(|py| -> Py<PyByteArray> {
PyByteArray::new(py, &serialized[..]).into()
});
Ok(b)
}
#[staticmethod]
pub fn from_bincode(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 ClassicalRegister")
})?,
})
}
pub fn to_json(&self) -> PyResult<String> {
serde_json::to_string(&self.internal)
.map_err(|_| PyRuntimeError::new_err("Unexpected error serializing ClassicalRegister"))
}
#[staticmethod]
pub fn from_json(json_string: &str) -> PyResult<Self> {
Ok(Self {
internal: serde_json::from_str(json_string).map_err(|_| {
PyValueError::new_err("Cannot deserialize string to ClassicalRegister")
})?,
})
}
pub fn __repr__(&self) -> String {
format!("{:?}", self.internal)
}
pub fn __copy__(&self) -> Self {
self.clone()
}
pub fn __deepcopy__(&self, _memodict: Py<PyAny>) -> Self {
self.clone()
}
fn __richcmp__(
&self,
other: ClassicalRegisterWrapper,
op: pyo3::class::basic::CompareOp,
) -> PyResult<bool> {
match op {
pyo3::class::basic::CompareOp::Eq => Ok(self.internal == other.internal),
pyo3::class::basic::CompareOp::Ne => Ok(self.internal != other.internal),
_ => Err(pyo3::exceptions::PyNotImplementedError::new_err(
"Other comparison not implemented",
)),
}
}
}
impl ClassicalRegisterWrapper {
pub fn from_pyany(input: Py<PyAny>) -> PyResult<ClassicalRegister> {
Python::with_gil(|py| -> PyResult<ClassicalRegister> {
let input = input.as_ref(py);
if let Ok(try_downcast) = input.extract::<ClassicalRegisterWrapper>() {
Ok(try_downcast.internal)
} else {
let get_bytes = input.call_method0("to_bincode").map_err(|_| {
PyTypeError::new_err("Python object cannot be converted to qoqo ClassicalRegister: Cast to binary representation failed".to_string())
})?;
let bytes = get_bytes.extract::<Vec<u8>>().map_err(|_| {
PyTypeError::new_err("Python object cannot be converted to qoqo ClassicalRegister: Cast to binary representation failed".to_string())
})?;
deserialize(&bytes[..]).map_err(|err| {
PyTypeError::new_err(format!(
"Python object cannot be converted to qoqo ClassicalRegister: Deserialization failed: {}",
err
))
})
}
})
}
}