use pyo3::{
prelude::*,
types::{PyDict, PyList},
};
use super::{to_pykey_err, to_pyvalue_err};
pub fn get_required_string(dict: &Bound<'_, PyDict>, key: &str) -> PyResult<String> {
dict.get_item(key)?
.ok_or_else(|| to_pykey_err(format!("Missing required key: {key}")))?
.extract()
}
pub fn get_required<T>(dict: &Bound<'_, PyDict>, key: &str) -> PyResult<T>
where
T: for<'a, 'py> FromPyObject<'a, 'py>,
for<'a, 'py> PyErr: From<<T as FromPyObject<'a, 'py>>::Error>,
{
dict.get_item(key)?
.ok_or_else(|| to_pykey_err(format!("Missing required key: {key}")))?
.extract()
.map_err(PyErr::from)
}
#[inline]
pub fn get_optional<T>(dict: &Bound<'_, PyDict>, key: &str) -> PyResult<Option<T>>
where
T: for<'a, 'py> FromPyObject<'a, 'py>,
for<'a, 'py> PyErr: From<<T as FromPyObject<'a, 'py>>::Error>,
{
match dict.get_item(key)? {
Some(value) => {
if value.is_none() {
Ok(None)
} else {
value.extract().map(Some).map_err(PyErr::from)
}
}
None => Ok(None),
}
}
pub fn get_required_parsed<T, F>(dict: &Bound<'_, PyDict>, key: &str, parser: F) -> PyResult<T>
where
F: FnOnce(String) -> Result<T, String>,
{
let value_str = get_required_string(dict, key)?;
parser(value_str).map_err(|e| to_pyvalue_err(format!("Failed to parse '{key}': {e}")))
}
pub fn get_optional_parsed<T, F>(
dict: &Bound<'_, PyDict>,
key: &str,
parser: F,
) -> PyResult<Option<T>>
where
F: FnOnce(String) -> Result<T, String>,
{
get_optional::<String>(dict, key)?
.map(parser)
.transpose()
.map_err(|e| to_pyvalue_err(format!("Failed to parse '{key}': {e}")))
}
pub fn get_required_list<'py>(
dict: &Bound<'py, PyDict>,
key: &str,
) -> PyResult<Bound<'py, PyList>> {
dict.get_item(key)?
.ok_or_else(|| to_pykey_err(format!("Missing required key: {key}")))?
.cast_into()
.map_err(Into::into)
}
#[cfg(test)]
mod tests {
use std::sync::Once;
use pyo3::exceptions::{PyKeyError, PyValueError};
use rstest::rstest;
use super::*;
fn ensure_python_initialized() {
static INIT: Once = Once::new();
INIT.call_once(Python::initialize);
}
#[rstest]
fn test_get_required_string() {
ensure_python_initialized();
Python::attach(|py| {
let dict = PyDict::new(py);
dict.set_item("name", "nautilus").unwrap();
let value = get_required_string(&dict, "name").unwrap();
let error = get_required_string(&dict, "missing").unwrap_err();
assert_eq!(value, "nautilus");
assert!(error.is_instance_of::<PyKeyError>(py));
assert_eq!(
error.value(py).to_string(),
"'Missing required key: missing'"
);
});
}
#[rstest]
fn test_get_optional_parsed() {
ensure_python_initialized();
Python::attach(|py| {
let dict = PyDict::new(py);
dict.set_item("value", "42").unwrap();
let parsed = get_optional_parsed(&dict, "value", |value| {
value.parse::<u64>().map_err(|e| e.to_string())
})
.unwrap();
let missing = get_optional_parsed(&dict, "missing", |value| {
value.parse::<u64>().map_err(|e| e.to_string())
})
.unwrap();
dict.set_item("value", py.None()).unwrap();
let none = get_optional_parsed(&dict, "value", |value| {
value.parse::<u64>().map_err(|e| e.to_string())
})
.unwrap();
dict.set_item("value", "invalid").unwrap();
let error =
get_optional_parsed::<u64, _>(&dict, "value", |_| Err("not a number".to_string()))
.unwrap_err();
assert_eq!(parsed, Some(42));
assert_eq!(missing, None);
assert_eq!(none, None);
assert!(error.is_instance_of::<PyValueError>(py));
assert_eq!(
error.value(py).to_string(),
"Failed to parse 'value': not a number"
);
});
}
}