use std::collections::{HashMap, HashSet};
use rill_core::math::Transcendental;
use rill_core::traits::{ParamValue, Params};
use crate::graph::GraphBuilder;
use serde::de;
use serde::{Deserialize, Serialize};
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct ResourceDef {
pub name: String,
pub kind: String,
pub capacity: usize,
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct GraphDef {
pub format_version: String,
pub sample_rate: f32,
pub block_size: usize,
#[serde(default)]
pub resources: Vec<ResourceDef>,
pub nodes: Vec<NodeDef>,
pub connections: Vec<ConnectionDef>,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub description: Option<String>,
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub enum NodeDef {
Source(SourceDef),
Processor(ProcessorDef),
Router(RouterDef),
Sink(SinkDef),
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct SourceDef {
pub id: u32,
pub type_name: String,
pub name: String,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub backend: Option<String>,
#[serde(
deserialize_with = "deserialize_params",
serialize_with = "serialize_params"
)]
pub parameters: HashMap<String, ParamValue>,
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct ProcessorDef {
pub id: u32,
pub type_name: String,
pub name: String,
#[serde(
deserialize_with = "deserialize_params",
serialize_with = "serialize_params"
)]
pub parameters: HashMap<String, ParamValue>,
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct RouterDef {
pub id: u32,
pub type_name: String,
pub name: String,
#[serde(
deserialize_with = "deserialize_params",
serialize_with = "serialize_params"
)]
pub parameters: HashMap<String, ParamValue>,
#[serde(default, skip_serializing_if = "Vec::is_empty")]
pub routing_matrix: Vec<RoutingEntry>,
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct RoutingEntry {
pub from: usize,
pub to: usize,
pub gain: f32,
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct SinkDef {
pub id: u32,
pub type_name: String,
pub name: String,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub backend: Option<String>,
#[serde(
deserialize_with = "deserialize_params",
serialize_with = "serialize_params"
)]
pub parameters: HashMap<String, ParamValue>,
}
impl NodeDef {
pub fn id(&self) -> u32 {
match self {
NodeDef::Source(s) => s.id,
NodeDef::Processor(p) => p.id,
NodeDef::Router(r) => r.id,
NodeDef::Sink(s) => s.id,
}
}
pub fn type_name(&self) -> &str {
match self {
NodeDef::Source(s) => &s.type_name,
NodeDef::Processor(p) => &p.type_name,
NodeDef::Router(r) => &r.type_name,
NodeDef::Sink(s) => &s.type_name,
}
}
pub fn name(&self) -> &str {
match self {
NodeDef::Source(s) => &s.name,
NodeDef::Processor(p) => &p.name,
NodeDef::Router(r) => &r.name,
NodeDef::Sink(s) => &s.name,
}
}
pub fn parameters(&self) -> &HashMap<String, ParamValue> {
match self {
NodeDef::Source(s) => &s.parameters,
NodeDef::Processor(p) => &p.parameters,
NodeDef::Router(r) => &r.parameters,
NodeDef::Sink(s) => &s.parameters,
}
}
pub fn parameters_mut(&mut self) -> &mut HashMap<String, ParamValue> {
match self {
NodeDef::Source(s) => &mut s.parameters,
NodeDef::Processor(p) => &mut p.parameters,
NodeDef::Router(r) => &mut r.parameters,
NodeDef::Sink(s) => &mut s.parameters,
}
}
pub fn backend(&self) -> Option<&str> {
match self {
NodeDef::Source(s) => s.backend.as_deref(),
NodeDef::Processor(_) => None,
NodeDef::Router(_) => None,
NodeDef::Sink(s) => s.backend.as_deref(),
}
}
pub fn kind_str(&self) -> &'static str {
match self {
NodeDef::Source(_) => "Source",
NodeDef::Processor(_) => "Processor",
NodeDef::Router(_) => "Router",
NodeDef::Sink(_) => "Sink",
}
}
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct ConnectionDef {
pub kind: SignalKind,
pub from_node: u32,
pub from_port: usize,
pub to_node: u32,
pub to_port: usize,
}
#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize)]
pub enum SignalKind {
Signal,
Control,
Clock,
Feedback,
}
#[derive(Debug, Clone)]
pub enum SerializationError {
UnknownType(String),
DuplicateNodeId(u32),
InvalidFormat(String),
}
impl std::fmt::Display for SerializationError {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
match self {
Self::UnknownType(t) => write!(f, "unknown node type: {t}"),
Self::DuplicateNodeId(id) => write!(f, "duplicate NodeId: {id}"),
Self::InvalidFormat(d) => write!(f, "invalid format: {d}"),
}
}
}
impl std::error::Error for SerializationError {}
impl GraphDef {
pub fn new(sample_rate: f32, block_size: usize) -> Self {
Self {
format_version: "rill/1".to_string(),
sample_rate,
block_size,
resources: Vec::new(),
nodes: Vec::new(),
connections: Vec::new(),
description: None,
}
}
pub fn add_node(&mut self, def: NodeDef) -> Result<(), SerializationError> {
let id = def.id();
if self.nodes.iter().any(|n| n.id() == id) {
return Err(SerializationError::DuplicateNodeId(id));
}
self.nodes.push(def);
Ok(())
}
pub fn add_connection(&mut self, conn: ConnectionDef) {
self.connections.push(conn);
}
pub fn set_node_param(&mut self, node_id: u32, key: &str, value: ParamValue) {
if let Some(nd) = self.nodes.iter_mut().find(|n| n.id() == node_id) {
nd.parameters_mut().insert(key.to_string(), value);
}
}
pub fn clear(&mut self) {
self.nodes.clear();
self.connections.clear();
}
}
impl GraphDef {
pub fn populate<T: Transcendental, const B: usize>(
&self,
builder: &mut GraphBuilder<T, B>,
) -> Result<(), SerializationError> {
builder.set_sample_rate(self.sample_rate);
let mut seen = HashSet::new();
for nd in &self.nodes {
if !seen.insert(nd.id()) {
return Err(SerializationError::DuplicateNodeId(nd.id()));
}
}
if self.block_size != B {
return Err(SerializationError::InvalidFormat(format!(
"expected block_size={B}, document has block_size={}",
self.block_size
)));
}
for rd in &self.resources {
builder.add_resource(crate::graph::GraphResource {
name: rd.name.clone(),
kind: rd.kind.clone(),
capacity: rd.capacity,
});
}
for nd in &self.nodes {
let mut p = Params::new(self.sample_rate);
for (k, v) in nd.parameters() {
p = p.with(k.clone(), v.clone());
}
let idx =
builder.add_node_with_name(nd.type_name(), &p, nd.id(), nd.name().to_string());
if let NodeDef::Router(ref r) = nd {
for entry in &r.routing_matrix {
builder.add_routing_entry(idx, entry.from, entry.to, entry.gain);
}
}
}
let id_to_idx: HashMap<u32, usize> = self
.nodes
.iter()
.enumerate()
.map(|(i, n)| (n.id(), i))
.collect();
for conn in &self.connections {
let from = *id_to_idx.get(&conn.from_node).ok_or_else(|| {
SerializationError::InvalidFormat(format!(
"connection references unknown from_node {}",
conn.from_node
))
})?;
let to = *id_to_idx.get(&conn.to_node).ok_or_else(|| {
SerializationError::InvalidFormat(format!(
"connection references unknown to_node {}",
conn.to_node
))
})?;
match conn.kind {
SignalKind::Signal => {
builder.connect_signal(from, conn.from_port, to, conn.to_port);
}
SignalKind::Control => {
builder.connect_control(from, conn.from_port, to, conn.to_port);
}
SignalKind::Clock => {
builder.connect_clock(from, conn.from_port, to, conn.to_port);
}
SignalKind::Feedback => {
builder.connect_feedback(from, conn.from_port, to, conn.to_port);
}
}
}
Ok(())
}
}
fn deserialize_params<'de, D>(deserializer: D) -> Result<HashMap<String, ParamValue>, D::Error>
where
D: serde::Deserializer<'de>,
{
let raw: HashMap<String, serde_json::Value> = HashMap::deserialize(deserializer)?;
raw.into_iter()
.map(|(k, v)| {
json_to_param_value(v)
.map(|pv| (k, pv))
.map_err(de::Error::custom)
})
.collect()
}
fn json_to_param_value(v: serde_json::Value) -> Result<ParamValue, String> {
match v {
serde_json::Value::Number(n) => n
.as_f64()
.map(|f| ParamValue::Float(f as f32))
.ok_or_else(|| "invalid number".to_string()),
serde_json::Value::String(s) => Ok(ParamValue::String(s)),
serde_json::Value::Bool(b) => Ok(ParamValue::Bool(b)),
serde_json::Value::Object(obj) => {
if let Some(val) = obj.get("Float").and_then(|v| v.as_f64()) {
return Ok(ParamValue::Float(val as f32));
}
if let Some(val) = obj.get("Int").and_then(|v| v.as_i64()) {
return Ok(ParamValue::Int(val as i32));
}
if let Some(val) = obj.get("Bool").and_then(|v| v.as_bool()) {
return Ok(ParamValue::Bool(val));
}
if let Some(val) = obj.get("String").and_then(|v| v.as_str()) {
return Ok(ParamValue::String(val.to_string()));
}
if let Some(val) = obj.get("Choice").and_then(|v| v.as_str()) {
return Ok(ParamValue::Choice(val.to_string()));
}
if let Some(arr) = obj.get("Bytes").and_then(|v| v.as_array()) {
let bytes: Vec<u8> = arr
.iter()
.filter_map(|v| v.as_u64().map(|n| n as u8))
.collect();
return Ok(ParamValue::Bytes(bytes));
}
Err("unknown variant in tagged format".to_string())
}
serde_json::Value::Array(arr) => {
let bytes: Vec<u8> = arr
.iter()
.filter_map(|v| v.as_u64().map(|n| n as u8))
.collect();
Ok(ParamValue::Bytes(bytes))
}
_ => Err("invalid param value type".to_string()),
}
}
fn serialize_params<S>(
params: &HashMap<String, ParamValue>,
serializer: S,
) -> Result<S::Ok, S::Error>
where
S: serde::Serializer,
{
use serde::ser::SerializeMap;
let mut map = serializer.serialize_map(Some(params.len()))?;
for (k, v) in params {
let json_val = param_value_to_json(v);
map.serialize_entry(k, &json_val)?;
}
map.end()
}
fn param_value_to_json(v: &ParamValue) -> serde_json::Value {
match v {
ParamValue::Float(f) => {
serde_json::Value::Number(serde_json::Number::from_f64(*f as f64).unwrap_or(0.into()))
}
ParamValue::Int(i) => serde_json::Value::Number(serde_json::Number::from(*i)),
ParamValue::Bool(b) => serde_json::Value::Bool(*b),
ParamValue::String(s) => serde_json::Value::String(s.clone()),
ParamValue::Choice(s) => serde_json::Value::String(s.clone()),
ParamValue::Bytes(b) => serde_json::Value::Array(
b.iter()
.map(|&x| serde_json::Value::Number(x.into()))
.collect(),
),
ParamValue::SignalSlab(_) => serde_json::Value::Null,
}
}
pub fn from_json(json: &str) -> Result<GraphDef, SerializationError> {
serde_json::from_str(json).map_err(|e| SerializationError::InvalidFormat(e.to_string()))
}
pub fn from_cbor(bytes: &[u8]) -> Result<GraphDef, SerializationError> {
serde_cbor::from_slice(bytes).map_err(|e| SerializationError::InvalidFormat(e.to_string()))
}