use super::Argument;
use super::NameOrParameter;
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())
}
}
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())
}
}
#[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));
}
}
}