use crate::CircuitWrapper;
use pyo3::exceptions::{PyRuntimeError, PyTypeError, PyValueError};
use pyo3::prelude::*;
use pyo3::types::PyByteArray;
use roqoqo::measurements::ClassicalRegister;
use roqoqo::prelude::*;
use roqoqo::Circuit;
#[cfg(feature = "json_schema")]
use roqoqo::ROQOQO_VERSION;
use std::collections::HashMap;
#[pyclass(
from_py_object,
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<&Bound<PyAny>>,
circuits: Vec<Py<PyAny>>,
) -> PyResult<Self> {
Python::attach(|py| -> PyResult<Self> {
let mut new_circuits: Vec<Circuit> = Vec::new();
for c in circuits.into_iter() {
let tmp_c = CircuitWrapper::from_pyany(c.bind(py)).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 = bincode::serde::encode_to_vec(&self.internal, bincode::config::legacy())
.map_err(|_| {
PyValueError::new_err("Cannot serialize ClassicalRegister to bytes")
})?;
let b: Py<PyByteArray> = Python::attach(|py| -> Py<PyByteArray> {
PyByteArray::new(py, &serialized[..]).into()
});
Ok(("ClassicalRegister", b))
}
pub fn to_bincode(&self) -> PyResult<Py<PyByteArray>> {
let serialized = bincode::serde::encode_to_vec(&self.internal, bincode::config::legacy())
.map_err(|_| {
PyValueError::new_err("Cannot serialize ClassicalRegister to bytes")
})?;
let b: Py<PyByteArray> = Python::attach(|py| -> Py<PyByteArray> {
PyByteArray::new(py, &serialized[..]).into()
});
Ok(b)
}
#[staticmethod]
pub fn from_bincode(input: &Bound<PyAny>) -> PyResult<Self> {
let bytes = input
.extract::<Vec<u8>>()
.map_err(|_| PyTypeError::new_err("Input cannot be converted to byte array"))?;
Ok(Self {
internal: bincode::serde::decode_from_slice(&bytes[..], bincode::config::legacy())
.map_err(|_| {
PyValueError::new_err("Input cannot be deserialized to ClassicalRegister")
})?
.0,
})
}
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: &Bound<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",
)),
}
}
#[cfg(feature = "json_schema")]
#[staticmethod]
pub fn json_schema() -> String {
let schema = schemars::schema_for!(ClassicalRegister);
serde_json::to_string_pretty(&schema).expect("Unexpected failure to serialize schema")
}
#[cfg(feature = "json_schema")]
#[staticmethod]
pub fn current_version() -> String {
ROQOQO_VERSION.to_string()
}
#[cfg(feature = "json_schema")]
pub fn min_supported_version(&self) -> String {
let min_version: (u32, u32, u32) =
ClassicalRegister::minimum_supported_roqoqo_version(&self.internal);
format!("{}.{}.{}", min_version.0, min_version.1, min_version.2)
}
}
impl ClassicalRegisterWrapper {
pub fn from_pyany(input: &Bound<PyAny>) -> PyResult<ClassicalRegister> {
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())
})?;
bincode::serde::decode_from_slice(&bytes[..], bincode::config::legacy()).map_err(|err| {
PyTypeError::new_err(format!(
"Python object cannot be converted to qoqo ClassicalRegister: Deserialization failed: {err}"
))
}).map(|(deserialized, _)| deserialized)
}
}
}