use super::Argument;
use super::NameOrParameter;
use crate::Result;
use qudit_core::ParamIndices;
use qudit_expr::Expression;
#[derive(Clone, Debug)]
pub struct ArgumentList {
entries: Vec<Argument>,
}
impl ArgumentList {
pub fn new(entries: Vec<Argument>) -> Self {
Self { entries }
}
pub fn arguments(&self) -> &[Argument] {
&self.entries
}
pub fn len(&self) -> usize {
self.entries.len()
}
pub fn is_empty(&self) -> bool {
self.entries.is_empty()
}
pub fn parameters(&self) -> Vec<NameOrParameter> {
let mut params = vec![];
for argument in self.entries.iter() {
for param in argument.parameters() {
if let NameOrParameter::Name(_) = param {
if params.contains(¶m) {
continue;
}
}
params.push(param);
}
}
params
}
pub fn variables(&self) -> Vec<String> {
let mut new_variables = vec![];
let mut unnamed_counter = 0;
for argument in self.entries.iter() {
new_variables.extend(argument.variables(&mut unnamed_counter))
}
new_variables
}
pub fn expressions(&self) -> Vec<Expression> {
let mut new_variables = vec![];
let mut unnamed_counter = 0;
for argument in self.entries.iter() {
new_variables.push(argument.as_substitution_expression(&mut unnamed_counter))
}
new_variables
}
pub fn requires_expression_modification(&self) -> bool {
self.entries
.iter()
.any(|arg| arg.requires_expression_modification())
}
pub fn slice_by_indices<I>(&self, indices: I) -> Self
where
I: IntoIterator<Item = usize>,
{
indices.into_iter().map(|i| self[i].clone()).collect()
}
pub fn map_indices_for_instruction<I>(&self, inner_indices: I) -> ParamIndices
where
I: IntoIterator<Item = usize>,
{
let unique_params = self.parameters();
let mut mapped_indices = Vec::new();
for i in inner_indices {
let argument = &self.entries[i];
for param in argument.parameters() {
if let Some(new_pos) = unique_params.iter().position(|p| p == ¶m) {
if !mapped_indices.contains(&new_pos) {
mapped_indices.push(new_pos);
}
}
}
}
ParamIndices::from(mapped_indices)
}
}
impl std::ops::Deref for ArgumentList {
type Target = [Argument];
fn deref(&self) -> &Self::Target {
&self.entries
}
}
impl<E: Into<Argument>> From<Vec<E>> for ArgumentList {
fn from(value: Vec<E>) -> Self {
Self::new(value.into_iter().map(|e| e.into()).collect())
}
}
impl<E: Into<Argument>, const N: usize> From<[E; N]> for ArgumentList {
fn from(value: [E; N]) -> Self {
Self::new(value.into_iter().map(|e| e.into()).collect())
}
}
impl<E: Into<Argument> + Clone, const N: usize> From<&[E; N]> for ArgumentList {
fn from(value: &[E; N]) -> Self {
Self::new(value.iter().map(|e| e.clone().into()).collect())
}
}
impl TryFrom<Vec<String>> for ArgumentList {
type Error = crate::Error;
fn try_from(value: Vec<String>) -> std::result::Result<Self, Self::Error> {
let mut arguments = Vec::with_capacity(value.len());
for s in value {
arguments.push(
Argument::try_from(s).map_err(|e| crate::Error::InvalidArgument {
message: e.to_string(),
})?,
);
}
Ok(ArgumentList::new(arguments))
}
}
impl TryFrom<Vec<&str>> for ArgumentList {
type Error = crate::Error;
fn try_from(value: Vec<&str>) -> std::result::Result<Self, Self::Error> {
let mut arguments = Vec::with_capacity(value.len());
for s in value {
arguments.push(
Argument::try_from(s).map_err(|e| crate::Error::InvalidArgument {
message: e.to_string(),
})?,
);
}
Ok(ArgumentList::new(arguments))
}
}
impl<const N: usize> TryFrom<[String; N]> for ArgumentList {
type Error = crate::Error;
fn try_from(value: [String; N]) -> std::result::Result<Self, Self::Error> {
let mut arguments = Vec::with_capacity(N);
for s in value {
arguments.push(
Argument::try_from(s).map_err(|e| crate::Error::InvalidArgument {
message: e.to_string(),
})?,
);
}
Ok(ArgumentList::new(arguments))
}
}
impl<const N: usize> TryFrom<[&str; N]> for ArgumentList {
type Error = crate::Error;
fn try_from(value: [&str; N]) -> std::result::Result<Self, Self::Error> {
let mut arguments = Vec::with_capacity(N);
for s in value {
arguments.push(
Argument::try_from(s).map_err(|e| crate::Error::InvalidArgument {
message: e.to_string(),
})?,
);
}
Ok(ArgumentList::new(arguments))
}
}
impl<E: Into<Argument>> FromIterator<E> for ArgumentList {
fn from_iter<T: IntoIterator<Item = E>>(iter: T) -> Self {
Self::new(iter.into_iter().map(|e| e.into()).collect())
}
}
pub trait IntoArgumentList {
fn into_args(self, num_args: usize) -> Result<ArgumentList>;
}
impl IntoArgumentList for ArgumentList {
fn into_args(self, num_args: usize) -> Result<ArgumentList> {
if num_args != self.len() {
Err(crate::Error::ArgumentListSizeMismatch {
actual: self.len() as u64,
expected: num_args as u64,
})
} else {
Ok(self)
}
}
}
impl IntoArgumentList for Option<ArgumentList> {
fn into_args(self, num_args: usize) -> Result<ArgumentList> {
match self {
Some(args) => args.into_args(num_args),
None => {
let args = ArgumentList::new(vec![Argument::Unspecified; num_args]);
Ok(args)
}
}
}
}
impl IntoArgumentList for () {
fn into_args(self, num_args: usize) -> Result<ArgumentList> {
let args = ArgumentList::new(vec![Argument::Unspecified; num_args]);
Ok(args)
}
}
impl<T> IntoArgumentList for T
where
T: TryInto<ArgumentList>,
T::Error: Into<crate::Error>,
{
fn into_args(self, num_args: usize) -> Result<ArgumentList> {
let args = self.try_into().map_err(Into::into)?;
args.into_args(num_args)
}
}
#[cfg(feature = "python")]
mod python_stub {
use super::ArgumentList;
use pyo3_stub_gen::impl_stub_type;
impl_stub_type!(ArgumentList = Vec<f64>);
}
#[cfg(feature = "python")]
mod python {
use super::Argument;
use super::ArgumentList;
use pyo3::exceptions::{PyTypeError, PyValueError};
use pyo3::prelude::*;
impl<'a, 'py> FromPyObject<'a, 'py> for ArgumentList {
type Error = PyErr;
fn extract(obj: Borrowed<'a, 'py, PyAny>) -> PyResult<Self> {
if let Ok(arg) = Argument::extract(obj) {
return Ok(ArgumentList::new(vec![arg]));
}
if let Ok(iter) = pyo3::types::PyIterator::from_object(&obj) {
let mut arguments = Vec::new();
for item in iter {
let item = item?;
match Argument::extract(item.as_borrowed()) {
Ok(arg) => arguments.push(arg),
Err(e) => {
return Err(PyValueError::new_err(format!(
"Cannot convert item to Argument: {}",
e
)));
}
}
}
return Ok(ArgumentList::new(arguments));
}
Err(PyTypeError::new_err(format!(
"Cannot convert {} to ArgumentList: must be iterable or convertible to Argument",
obj.get_type().name()?,
)))
}
}
#[cfg(test)]
mod tests {
use super::*;
use pyo3::Python;
use pyo3::types::{PyDict, PyFloat, PyList, PyString, PyTuple};
#[test]
fn test_from_py_empty_list() {
Python::initialize();
Python::attach(|py| {
let empty_list = PyList::empty(py);
let result = ArgumentList::extract(empty_list.as_any().as_borrowed()).unwrap();
assert_eq!(result.len(), 0);
assert!(result.is_empty());
});
}
#[test]
fn test_from_py_list_of_numbers() {
Python::initialize();
Python::attach(|py| {
let py_list = PyList::new(py, [1.0, 2.5, 3.14]).unwrap();
let result = ArgumentList::extract(py_list.as_any().as_borrowed()).unwrap();
assert_eq!(result.len(), 3);
let args = result.arguments();
assert!(matches!(args[0], Argument::Float64(1.0)));
assert!(matches!(args[1], Argument::Float64(2.5)));
assert!(matches!(args[2], Argument::Float64(3.14)));
});
}
#[test]
fn test_from_py_tuple_mixed_types() {
Python::initialize();
Python::attach(|py| {
let one = PyFloat::new(py, 1.5);
let two = PyString::new(py, "x + 1");
let none = py.None().into_bound(py);
let items: Vec<&Bound<'_, PyAny>> = vec![one.as_any(), two.as_any(), &none];
let py_tuple = PyTuple::new(py, items).unwrap();
let result = ArgumentList::extract(py_tuple.as_any().as_borrowed()).unwrap();
assert_eq!(result.len(), 3);
let args = result.arguments();
assert!(matches!(args[0], Argument::Float64(1.5)));
assert!(matches!(args[1], Argument::Expression(_)));
assert!(matches!(args[2], Argument::Unspecified));
});
}
#[test]
fn test_from_py_single_argument() {
Python::initialize();
Python::attach(|py| {
let float_val = PyFloat::new(py, 42.0);
let result = ArgumentList::extract(float_val.as_any().as_borrowed()).unwrap();
assert_eq!(result.len(), 1);
let args = result.arguments();
assert!(matches!(args[0], Argument::Float64(42.0)));
});
}
#[test]
fn test_from_py_list_with_expressions() {
Python::initialize();
Python::attach(|py| {
let one = PyString::new(py, "sin(x)");
let two = PyString::new(py, "a*b + c");
let three = PyString::new(py, "pi/4");
let expressions = vec![one.as_any(), two.as_any(), three.as_any()];
let py_list = PyList::new(py, expressions).unwrap();
let result = ArgumentList::extract(py_list.as_any().as_borrowed()).unwrap();
assert_eq!(result.len(), 3);
let args = result.arguments();
for arg in args {
assert!(matches!(arg, Argument::Expression(_)));
}
});
}
#[test]
fn test_from_py_list_with_invalid_item() {
Python::initialize();
Python::attach(|py| {
let one = PyFloat::new(py, 1.0);
let two = PyDict::new(py); let items: Vec<&Bound<'_, PyAny>> = vec![one.as_any(), two.as_any()];
let py_list = PyList::new(py, items).unwrap();
let result = ArgumentList::extract(py_list.as_any().as_borrowed());
assert!(result.is_err());
let err = result.unwrap_err();
assert!(err.is_instance_of::<pyo3::exceptions::PyValueError>(py));
assert!(err.to_string().contains("Cannot convert item to Argument"));
});
}
#[test]
fn test_from_py_list_with_complex_expression_rejected() {
Python::initialize();
Python::attach(|py| {
let complex_expr = PyString::new(py, "1 + i");
let items: Vec<&Bound<'_, PyAny>> = vec![
complex_expr.as_any(), ];
let py_list = PyList::new(py, items).unwrap();
let result = ArgumentList::extract(py_list.as_any().as_borrowed());
assert!(result.is_err());
let err = result.unwrap_err();
assert!(err.is_instance_of::<pyo3::exceptions::PyValueError>(py));
});
}
#[test]
fn test_from_py_list_with_unnamed_rejected() {
Python::initialize();
Python::attach(|py| {
let unnamed_expr = PyString::new(py, "unnamed_var");
let items: Vec<&Bound<'_, PyAny>> = vec![unnamed_expr.as_any()];
let py_list = PyList::new(py, items).unwrap();
let result = ArgumentList::extract(py_list.as_any().as_borrowed());
assert!(result.is_err());
let err = result.unwrap_err();
assert!(err.is_instance_of::<pyo3::exceptions::PyValueError>(py));
assert!(
err.to_string()
.contains("Expression arguments cannot contain 'unnamed_'")
);
});
}
#[test]
fn test_from_py_single_string_expression() {
Python::initialize();
Python::attach(|py| {
let string_val = PyString::new(py, "x*y + z");
let result = ArgumentList::extract(string_val.as_any().as_borrowed()).unwrap();
assert_eq!(result.len(), 1);
let args = result.arguments();
assert!(matches!(args[0], Argument::Expression(_)));
});
}
#[test]
fn test_parameters_extraction() {
Python::initialize();
Python::attach(|py| {
let one = PyString::new(py, "x");
let two = PyString::new(py, "x + y"); let three = PyFloat::new(py, 1.0);
let expressions = vec![one.as_any(), two.as_any(), three.as_any()];
let py_list = PyList::new(py, expressions).unwrap();
let result = ArgumentList::extract(py_list.as_any().as_borrowed()).unwrap();
let params = result.parameters();
assert_eq!(params.len(), 3);
});
}
}
}
#[cfg(test)]
mod tests {
use super::super::Parameter;
use super::*;
use qudit_expr::Expression;
#[test]
fn test_argumentlist_new() {
let args = vec![Argument::Float64(1.0), Argument::Float64(2.0)];
let list = ArgumentList::new(args);
assert_eq!(list.len(), 2);
assert!(!list.is_empty());
}
#[test]
fn test_argumentlist_empty() {
let list = ArgumentList::new(vec![]);
assert_eq!(list.len(), 0);
assert!(list.is_empty());
}
#[test]
fn test_argumentlist_arguments_access() {
let args = vec![Argument::Float64(1.0), Argument::Unspecified];
let list = ArgumentList::new(args);
let accessed_args = list.arguments();
assert_eq!(accessed_args.len(), 2);
assert!(matches!(accessed_args[0], Argument::Float64(1.0)));
assert!(matches!(accessed_args[1], Argument::Unspecified));
}
#[test]
fn test_parameters_single_argument() {
let args = vec![Argument::Float64(42.0)];
let list = ArgumentList::new(args);
let params = list.parameters();
assert_eq!(params.len(), 1);
assert!(matches!(
params[0],
NameOrParameter::Parameter(Parameter::Assigned64(42.0))
));
}
#[test]
fn test_parameters_deduplication() {
let x_expr = Argument::try_from("x").unwrap();
let xy_expr = Argument::try_from("x + y").unwrap();
let args = vec![x_expr, xy_expr];
let list = ArgumentList::new(args);
let params = list.parameters();
assert_eq!(params.len(), 2);
let param_names: Vec<String> = params
.into_iter()
.map(|p| match p {
NameOrParameter::Name(name) => name,
_ => panic!("Expected Named parameter"),
})
.collect();
assert!(param_names.contains(&"x".to_string()));
assert!(param_names.contains(&"y".to_string()));
}
#[test]
fn test_parameters_mixed_types() {
let args = vec![
Argument::Float32(1.0),
Argument::Float64(2.0),
Argument::Unspecified,
Argument::try_from("theta").unwrap(),
];
let list = ArgumentList::new(args);
let params = list.parameters();
assert_eq!(params.len(), 4);
assert!(matches!(
params[0],
NameOrParameter::Parameter(Parameter::Assigned32(1.0))
));
assert!(matches!(
params[1],
NameOrParameter::Parameter(Parameter::Assigned64(2.0))
));
assert!(matches!(
params[2],
NameOrParameter::Parameter(Parameter::Unassigned)
));
assert!(matches!(params[3], NameOrParameter::Name(ref name) if name == "theta"));
}
#[test]
fn test_variables_generation() {
let args = vec![
Argument::Float64(1.0), Argument::try_from("x").unwrap(), Argument::Unspecified, ];
let list = ArgumentList::new(args);
let vars = list.variables();
assert_eq!(vars.len(), 3);
assert_eq!(vars[0], "unnamed_0");
assert_eq!(vars[1], "x");
assert_eq!(vars[2], "unnamed_1");
}
#[test]
fn test_variables_constant_expression() {
let constant_expr = Expression::from_float_64(3.14);
let args = vec![Argument::Expression(constant_expr), Argument::Float64(2.0)];
let list = ArgumentList::new(args);
let vars = list.variables();
assert_eq!(vars.len(), 2);
assert_eq!(vars[0], "unnamed_0"); assert_eq!(vars[1], "unnamed_1"); }
#[test]
fn test_expressions_generation() {
let args = vec![Argument::Float64(1.0), Argument::try_from("x + 1").unwrap()];
let list = ArgumentList::new(args.clone());
let exprs = list.expressions();
assert_eq!(exprs.len(), 2);
assert_eq!(exprs[0], Expression::Variable("unnamed_0".to_string()));
if let Argument::Expression(original_expr) = &args[1] {
assert_eq!(exprs[1], *original_expr);
} else {
panic!("Expected Expression argument");
}
}
#[test]
fn test_from_vec() {
let values = vec![1.0f64, 2.0f64, 3.0f64];
let list = ArgumentList::from(values);
assert_eq!(list.len(), 3);
let args = list.arguments();
assert!(matches!(args[0], Argument::Float64(1.0)));
assert!(matches!(args[1], Argument::Float64(2.0)));
assert!(matches!(args[2], Argument::Float64(3.0)));
}
#[test]
fn test_from_array() {
let values = [1.0f32, 2.0f32];
let list = ArgumentList::from(values);
assert_eq!(list.len(), 2);
let args = list.arguments();
assert!(matches!(args[0], Argument::Float32(1.0)));
assert!(matches!(args[1], Argument::Float32(2.0)));
}
#[test]
fn test_from_array_ref() {
let values = [1.0f64, 2.0f64];
let list = ArgumentList::from(&values);
assert_eq!(list.len(), 2);
let args = list.arguments();
assert!(matches!(args[0], Argument::Float64(1.0)));
assert!(matches!(args[1], Argument::Float64(2.0)));
}
#[test]
fn test_argumentlist_clone() {
let args = vec![Argument::Float64(42.0), Argument::Unspecified];
let list1 = ArgumentList::new(args);
let list2 = list1.clone();
assert_eq!(list1.len(), list2.len());
assert_eq!(list1.arguments().len(), list2.arguments().len());
}
#[test]
fn test_argumentlist_debug() {
let args = vec![Argument::Float64(1.0)];
let list = ArgumentList::new(args);
let debug_str = format!("{:?}", list);
assert!(debug_str.contains("ArgumentList"));
assert!(debug_str.contains("Float64"));
}
#[test]
fn test_complex_multivariate_parameter_extraction() {
let args = vec![
Argument::try_from("a*b").unwrap(), Argument::try_from("c + a").unwrap(), Argument::Float64(1.0), ];
let list = ArgumentList::new(args);
let params = list.parameters();
assert_eq!(params.len(), 4);
let named_count = params
.iter()
.filter(|p| matches!(p, NameOrParameter::Name(_)))
.count();
let constant_count = params
.iter()
.filter(|p| matches!(p, NameOrParameter::Parameter(Parameter::Assigned64(_))))
.count();
assert_eq!(named_count, 3); assert_eq!(constant_count, 1); }
#[test]
fn test_large_argument_list() {
let mut args = Vec::new();
for i in 0..100 {
args.push(Argument::Float64(i as f64));
}
let list = ArgumentList::new(args);
assert_eq!(list.len(), 100);
assert!(!list.is_empty());
let params = list.parameters();
assert_eq!(params.len(), 100);
let vars = list.variables();
assert_eq!(vars.len(), 100); for (i, var) in vars.iter().enumerate() {
assert_eq!(*var, format!("unnamed_{}", i));
}
}
}
#[test]
fn test_from_iterator() {
let values = vec![1.0f64, 2.0f64, 3.0f64];
let list: ArgumentList = values.into_iter().collect();
assert_eq!(list.len(), 3);
let args = list.arguments();
assert!(matches!(args[0], Argument::Float64(1.0)));
assert!(matches!(args[1], Argument::Float64(2.0)));
assert!(matches!(args[2], Argument::Float64(3.0)));
}
#[test]
fn test_collect_from_range() {
let list: ArgumentList = (0..5).map(|i| i as f64).collect();
assert_eq!(list.len(), 5);
let args = list.arguments();
for (i, arg) in args.iter().enumerate() {
assert!(matches!(arg, Argument::Float64(val) if *val == i as f64));
}
}
#[test]
fn test_collect_mixed_types() {
let mixed_args: Vec<Argument> = vec![
Argument::Float64(1.0),
Argument::Float32(2.0),
Argument::Unspecified,
];
let list: ArgumentList = mixed_args.into_iter().collect();
assert_eq!(list.len(), 3);
let args = list.arguments();
assert!(matches!(args[0], Argument::Float64(1.0)));
assert!(matches!(args[1], Argument::Float32(2.0)));
assert!(matches!(args[2], Argument::Unspecified));
}
#[test]
fn test_try_from_vec_string() {
let strings = vec!["x".to_string(), "y + 1".to_string(), "pi/2".to_string()];
let result = ArgumentList::try_from(strings);
assert!(result.is_ok());
let list = result.unwrap();
assert_eq!(list.len(), 3);
let args = list.arguments();
for arg in args {
assert!(matches!(arg, Argument::Expression(_)));
}
}
#[test]
fn test_try_from_vec_str() {
let strings = vec!["a", "b*c", "sin(theta)"];
let result = ArgumentList::try_from(strings);
assert!(result.is_ok());
let list = result.unwrap();
assert_eq!(list.len(), 3);
let args = list.arguments();
for arg in args {
assert!(matches!(arg, Argument::Expression(_)));
}
}
#[test]
fn test_try_from_array_str() {
let strings = ["x", "y"];
let result = ArgumentList::try_from(strings);
assert!(result.is_ok());
let list = result.unwrap();
assert_eq!(list.len(), 2);
}
#[test]
fn test_into_argument_list_from_vec_strings() {
let strings = vec!["x", "y"];
let result = strings.into_args(2);
assert!(result.is_ok());
let list = result.unwrap();
assert_eq!(list.len(), 2);
let args = list.arguments();
for arg in args {
assert!(matches!(arg, Argument::Expression(_)));
}
}
#[test]
fn test_into_argument_list_from_vec_owned_strings() {
let strings = vec!["x + 1".to_string(), "y * 2".to_string()];
let result = strings.into_args(2);
assert!(result.is_ok());
let list = result.unwrap();
assert_eq!(list.len(), 2);
}
#[test]
fn test_into_argument_list_size_mismatch() {
let strings = vec!["x", "y"];
let result = strings.into_args(3); assert!(result.is_err());
}