use crate::error::{HessboostError, Result};
use serde_json::Value;
pub(super) fn field<'a>(v: &'a Value, key: &str) -> Result<&'a Value> {
v.get(key).ok_or_else(|| HessboostError::missing_field(key))
}
pub(super) fn optional_str<'a>(v: &'a Value, key: &str) -> Result<Option<&'a str>> {
v.get(key)
.map(|value| {
value
.as_str()
.ok_or_else(|| HessboostError::model_format(format!("invalid `{key}` {value}")))
})
.transpose()
}
pub(super) fn count_param(v: &Value, key: &str, default: usize) -> Result<usize> {
let Some(value) = v.get(key) else {
return Ok(default);
};
scalar_count(value)
.ok_or_else(|| HessboostError::model_format(format!("invalid `{key}` {value}")))
}
pub(super) fn scalar_count(value: &Value) -> Option<usize> {
const EXACT: f64 = 9_007_199_254_740_992.0;
let integer = match value {
Value::Number(n) => n.as_u64(),
Value::String(s) => s.parse::<u64>().ok(),
_ => None,
};
integer
.or_else(|| {
scalar_f64(value)
.filter(|&n| (0.0..=EXACT).contains(&n) && n.fract() == 0.0)
.map(|n| n as u64)
})
.and_then(|n| usize::try_from(n).ok())
}
pub(super) fn scalar_f64(v: &Value) -> Option<f64> {
match v {
Value::Number(n) => n.as_f64(),
Value::String(s) => s.parse::<f64>().ok(),
Value::Bool(b) => Some(if *b { 1.0 } else { 0.0 }),
_ => None,
}
}
fn column<'a>(v: &'a Value, key: &str, len: Option<usize>) -> Result<&'a [Value]> {
let entries = v
.get(key)
.ok_or_else(|| HessboostError::missing_field(key))?
.as_array()
.ok_or_else(|| HessboostError::model_format(format!("`{key}` is not an array")))?;
match len {
Some(len) if entries.len() != len => Err(HessboostError::model_format(format!(
"`{key}` has {} entries, not {len}",
entries.len()
))),
_ => Ok(entries),
}
}
pub(super) fn column_len(v: &Value, key: &str) -> Result<usize> {
column(v, key, None).map(<[Value]>::len)
}
pub(super) fn float_column(v: &Value, key: &str, len: Option<usize>) -> Result<Vec<f32>> {
column(v, key, len)?
.iter()
.map(|entry| {
scalar_f64(entry).map(|x| x as f32).ok_or_else(|| {
HessboostError::model_format(format!("`{key}` contains a non-numeric entry"))
})
})
.collect()
}
pub(super) fn optional_float_column(v: &Value, key: &str, len: usize) -> Result<Option<Vec<f32>>> {
v.get(key)
.map(|_| float_column(v, key, Some(len)))
.transpose()
}
pub(super) fn integer_column<T: TryFrom<i64>>(
v: &Value,
key: &str,
len: usize,
range: std::ops::RangeInclusive<i64>,
) -> Result<Vec<T>> {
column(v, key, Some(len))?
.iter()
.map(|entry| {
let value = scalar_f64(entry).ok_or_else(|| {
HessboostError::model_format(format!("`{key}` contains a non-numeric entry"))
})?;
let integer = value as i64;
if value.fract() != 0.0 || !range.contains(&integer) || integer as f64 != value {
return Err(HessboostError::model_format(format!(
"`{key}` contains an invalid entry {value}"
)));
}
T::try_from(integer).map_err(|_| {
HessboostError::model_format(format!("`{key}` contains an invalid entry {value}"))
})
})
.collect()
}
pub(super) fn strict_nonnegative_integer_array(v: &Value, key: &str) -> Result<Vec<u64>> {
let Some(value) = v.get(key) else {
return Ok(Vec::new());
};
let entries = value
.as_array()
.ok_or_else(|| HessboostError::model_format(format!("`{key}` is not an array")))?;
entries
.iter()
.map(|entry| {
let value = scalar_f64(entry).ok_or_else(|| {
HessboostError::model_format(format!("`{key}` contains a non-numeric entry"))
})?;
if !value.is_finite() || value < 0.0 || value.fract() != 0.0 || value > u64::MAX as f64
{
return Err(HessboostError::model_format(format!(
"`{key}` contains an invalid integer {value}"
)));
}
Ok(value as u64)
})
.collect()
}