use serde::{Deserialize, Deserializer};
use super::error::JsonError;
const MAX_WEIGHTS: usize = (256 * 1024 * 1024 / 4) as usize;
pub const MAX_LAYERS: usize = 8;
pub const MAX_HIDDEN_SIZE: usize = 512;
const MAX_TRAINING_BYTES: usize = 1024 * 1024;
const MAX_TRAINING_DEPTH: usize = 16;
const MAX_SUBMODELS: usize = 8;
pub const MAX_LSTM_LAYERS: usize = 16;
pub const MAX_LSTM_HIDDEN_SIZE: usize = 1024;
pub const MAX_WAVENET_FREE_CHANNELS: usize = 512;
pub const MAX_A2_DYN_CHANNELS: usize = 256;
pub const MAX_A2_DYN_BOTTLENECK: usize = 256;
pub const MAX_KERNEL_SIZE: usize = 64;
pub const MAX_DILATION: usize = 4096;
pub const MAX_DILATIONS_PER_ARRAY: usize = 64;
pub const MAX_WAVENET_ARRAYS: usize = 8;
pub const MAX_HEAD_SIZE: usize = 512;
pub const MAX_CONVNET_CHANNELS: usize = 512;
pub const MAX_CONVNET_KERNEL_SIZE: usize = 64;
pub const MAX_RECEPTIVE_FIELD: usize = 65536;
pub const MAX_TOTAL_STATE_FRAMES: usize = 1 << 26;
#[cfg(not(test))]
struct WeightsVisitor;
#[cfg(test)]
pub(super) struct WeightsVisitor;
impl<'de> serde::de::Visitor<'de> for WeightsVisitor {
type Value = Vec<f32>;
fn expecting(&self, formatter: &mut std::fmt::Formatter) -> std::fmt::Result {
formatter.write_str("a sequence of f32 floats within the size limit")
}
fn visit_seq<A>(self, mut seq: A) -> Result<Vec<f32>, A::Error>
where
A: serde::de::SeqAccess<'de>,
{
let mut weights = Vec::new();
let mut index: usize = 0;
loop {
match seq.next_element::<f32>() {
Ok(Some(val)) => {
if !val.is_finite() {
return Err(serde::de::Error::custom(JsonError::WeightNotFinite {
index,
value: val,
}));
}
if weights.len() >= MAX_WEIGHTS {
return Err(serde::de::Error::custom(JsonError::WeightsExceedLimit {
got: weights.len() + 1,
max: MAX_WEIGHTS,
}));
}
weights.push(val);
index += 1;
}
Ok(None) => break,
Err(e) => return Err(e),
}
}
Ok(weights)
}
}
pub(crate) fn deserialize_weights<'de, D>(deserializer: D) -> Result<Vec<f32>, D::Error>
where
D: Deserializer<'de>,
{
deserializer.deserialize_seq(WeightsVisitor)
}
struct LimitedValueVisitor {
depth: usize,
max_depth: usize,
max_size: usize,
current_size: std::cell::Cell<usize>,
}
impl LimitedValueVisitor {
fn root(max_depth: usize, max_size: usize) -> Self {
Self {
depth: 0,
max_depth,
max_size,
current_size: std::cell::Cell::new(0),
}
}
fn child(&self) -> Self {
Self {
depth: self.depth + 1,
max_depth: self.max_depth,
max_size: self.max_size,
current_size: self.current_size.clone(),
}
}
fn add_size(&self, bytes: usize) -> Result<(), serde_json::Error> {
let new = self.current_size.get() + bytes;
self.current_size.set(new);
if new > self.max_size {
return Err(serde::de::Error::custom(JsonError::TrainingTooLarge {
size: new,
max_size: self.max_size,
}));
}
Ok(())
}
fn check_depth(&self) -> Result<(), serde_json::Error> {
if self.depth > self.max_depth {
return Err(serde::de::Error::custom(JsonError::TrainingTooDeep {
depth: self.depth,
max_depth: self.max_depth,
}));
}
Ok(())
}
}
impl<'de> serde::de::Visitor<'de> for LimitedValueVisitor {
type Value = serde_json::Value;
fn expecting(&self, formatter: &mut std::fmt::Formatter) -> std::fmt::Result {
formatter.write_str("a JSON value within depth and size limits")
}
fn visit_bool<E>(self, v: bool) -> Result<serde_json::Value, E>
where
E: serde::de::Error,
{
self.add_size(if v { 4 } else { 5 }).map_err(E::custom)?;
Ok(serde_json::Value::Bool(v))
}
fn visit_i64<E>(self, v: i64) -> Result<serde_json::Value, E>
where
E: serde::de::Error,
{
self.add_size(16).map_err(E::custom)?;
Ok(serde_json::Value::Number(serde_json::Number::from(v)))
}
fn visit_u64<E>(self, v: u64) -> Result<serde_json::Value, E>
where
E: serde::de::Error,
{
self.add_size(16).map_err(E::custom)?;
Ok(serde_json::Value::Number(serde_json::Number::from(v)))
}
fn visit_f64<E>(self, v: f64) -> Result<serde_json::Value, E>
where
E: serde::de::Error,
{
self.add_size(16).map_err(E::custom)?;
Ok(serde_json::Value::Number(
serde_json::Number::from_f64(v).unwrap_or(serde_json::Number::from(0)),
))
}
fn visit_str<E>(self, v: &str) -> Result<serde_json::Value, E>
where
E: serde::de::Error,
{
self.add_size(v.len() + 2).map_err(E::custom)?;
Ok(serde_json::Value::String(v.to_string()))
}
fn visit_string<E>(self, v: String) -> Result<serde_json::Value, E>
where
E: serde::de::Error,
{
self.add_size(v.len() + 2).map_err(E::custom)?;
Ok(serde_json::Value::String(v))
}
fn visit_unit<E>(self) -> Result<serde_json::Value, E> {
Ok(serde_json::Value::Null)
}
fn visit_seq<A>(self, mut seq: A) -> Result<serde_json::Value, A::Error>
where
A: serde::de::SeqAccess<'de>,
{
self.check_depth().map_err(serde::de::Error::custom)?;
self.add_size(2).map_err(serde::de::Error::custom)?;
let mut arr = Vec::new();
loop {
match seq.next_element_seed(self.child()) {
Ok(Some(val)) => {
self.add_size(1).map_err(serde::de::Error::custom)?;
arr.push(val);
}
Ok(None) => break,
Err(e) => return Err(e),
}
}
Ok(serde_json::Value::Array(arr))
}
fn visit_map<A>(self, mut map: A) -> Result<serde_json::Value, A::Error>
where
A: serde::de::MapAccess<'de>,
{
self.check_depth().map_err(serde::de::Error::custom)?;
self.add_size(2).map_err(serde::de::Error::custom)?;
let mut obj = serde_json::Map::new();
loop {
match map.next_key::<String>() {
Ok(Some(key)) => {
let key_len = key.len() + 4;
self.add_size(key_len).map_err(serde::de::Error::custom)?;
let val: serde_json::Value = map.next_value_seed(self.child())?;
obj.insert(key, val);
self.add_size(1).map_err(serde::de::Error::custom)?;
}
Ok(None) => break,
Err(e) => return Err(e),
}
}
Ok(serde_json::Value::Object(obj))
}
}
impl<'de> serde::de::DeserializeSeed<'de> for LimitedValueVisitor {
type Value = serde_json::Value;
fn deserialize<D>(self, deserializer: D) -> Result<Self::Value, D::Error>
where
D: serde::Deserializer<'de>,
{
deserializer.deserialize_any(self)
}
}
struct TrainingOptionVisitor;
impl<'de> serde::de::Visitor<'de> for TrainingOptionVisitor {
type Value = Option<serde_json::Value>;
fn expecting(&self, formatter: &mut std::fmt::Formatter) -> std::fmt::Result {
formatter.write_str("an optional JSON value")
}
fn visit_none<E>(self) -> Result<Self::Value, E>
where
E: serde::de::Error,
{
Ok(None)
}
fn visit_unit<E>(self) -> Result<Self::Value, E>
where
E: serde::de::Error,
{
Ok(None)
}
fn visit_some<D>(self, deserializer: D) -> Result<Self::Value, D::Error>
where
D: serde::Deserializer<'de>,
{
let visitor = LimitedValueVisitor::root(MAX_TRAINING_DEPTH, MAX_TRAINING_BYTES);
let value: serde_json::Value = deserializer.deserialize_any(visitor)?;
Ok(Some(value))
}
}
pub(crate) fn deserialize_training<'de, D>(
deserializer: D,
) -> Result<Option<serde_json::Value>, D::Error>
where
D: Deserializer<'de>,
{
deserializer.deserialize_option(TrainingOptionVisitor)
}
pub(crate) fn deserialize_submodels<'de, D>(
deserializer: D,
) -> Result<Option<Vec<serde_json::Value>>, D::Error>
where
D: Deserializer<'de>,
{
deserializer.deserialize_option(SubmodelsOptionVisitor)
}
struct SubmodelsOptionVisitor;
impl<'de> serde::de::Visitor<'de> for SubmodelsOptionVisitor {
type Value = Option<Vec<serde_json::Value>>;
fn expecting(&self, formatter: &mut std::fmt::Formatter) -> std::fmt::Result {
formatter.write_str("an optional array of submodel entries")
}
fn visit_none<E>(self) -> Result<Self::Value, E>
where
E: serde::de::Error,
{
Ok(None)
}
fn visit_unit<E>(self) -> Result<Self::Value, E>
where
E: serde::de::Error,
{
Ok(None)
}
fn visit_some<D>(self, deserializer: D) -> Result<Self::Value, D::Error>
where
D: serde::Deserializer<'de>,
{
let arr: Vec<serde_json::Value> = Vec::deserialize(deserializer)?;
if arr.len() > MAX_SUBMODELS {
return Err(serde::de::Error::custom(JsonError::SubmodelsExceedLimit {
got: arr.len(),
max: MAX_SUBMODELS,
}));
}
Ok(Some(arr))
}
}
pub(crate) fn deserialize_sample_rate<'de, D>(deserializer: D) -> Result<Option<f32>, D::Error>
where
D: Deserializer<'de>,
{
struct SampleRateOptionVisitor;
impl<'de> serde::de::Visitor<'de> for SampleRateOptionVisitor {
type Value = Option<f32>;
fn expecting(&self, formatter: &mut std::fmt::Formatter) -> std::fmt::Result {
formatter.write_str("an optional f32 sample rate")
}
fn visit_none<E>(self) -> Result<Self::Value, E>
where
E: serde::de::Error,
{
Ok(None)
}
fn visit_unit<E>(self) -> Result<Self::Value, E>
where
E: serde::de::Error,
{
Ok(None)
}
fn visit_some<D>(self, deserializer: D) -> Result<Self::Value, D::Error>
where
D: serde::Deserializer<'de>,
{
let val = f32::deserialize(deserializer)?;
if !val.is_finite() {
return Err(serde::de::Error::custom(JsonError::InvalidSampleRate {
value: val,
reason: "must be finite",
}));
}
if val <= 0.0 {
return Ok(None);
}
Ok(Some(val))
}
}
deserializer.deserialize_option(SampleRateOptionVisitor)
}