use anyhow::{Context, Result};
use midir::os::unix::{VirtualInput, VirtualOutput};
use midir::{MidiInput, MidiInputConnection, MidiOutput, MidiOutputConnection};
use mlua::prelude::*;
use mlua::LuaSerdeExt;
use std::collections::HashMap;
use std::net::SocketAddr;
use std::path::Path;
use std::rc::Rc;
use std::sync::{Arc, Mutex};
use tokio::sync::mpsc;
use tracing::{error, info, warn};
#[derive(Debug, Clone, Default)]
pub struct ConnectDecl {
pub inputs: HashMap<String, Vec<String>>,
pub outputs: HashMap<String, Vec<String>>,
}
#[derive(Debug, Clone, Default, PartialEq)]
pub struct ObsDecl {
pub connections: Vec<String>,
}
use crate::config::Config;
use crate::lua_api::{
lua_to_midi_bytes, lua_val_to_osc_type, midi_bytes_to_lua, osc_message_to_lua,
toml_table_to_lua,
};
use crate::osc::{OscDecl, OscSender};
use crate::timer::{Timer, TimerEvent};
enum RouteEvent {
Midi { port: String, bytes: Vec<u8> },
Timer(TimerEvent),
Osc { from: SocketAddr, address: String, args: Vec<rosc::OscType> },
ObsEvent { connection: String, event: serde_json::Value },
ResyncState,
Shutdown,
}
#[derive(Debug, Clone, PartialEq)]
pub struct PortDecl {
pub inputs: Vec<String>,
pub outputs: Vec<String>,
}
impl Default for PortDecl {
fn default() -> Self {
PortDecl {
inputs: vec!["default".to_string()],
outputs: vec!["default".to_string()],
}
}
}
impl PortDecl {
pub fn is_default(&self) -> bool {
self.inputs.len() == 1
&& self.inputs[0] == "default"
&& self.outputs.len() == 1
&& self.outputs[0] == "default"
}
}
pub struct RoutePorts {
out_conns: HashMap<String, Arc<Mutex<MidiOutputConnection>>>,
midi_fwds: HashMap<String, Arc<Mutex<Option<mpsc::Sender<RouteEvent>>>>>,
_in_conns: Vec<MidiInputConnection<()>>,
pub decl: PortDecl,
}
impl RoutePorts {
fn create(
route_name: &str,
decl: &PortDecl,
initial_tx: &mpsc::Sender<RouteEvent>,
) -> Result<Rc<Self>> {
let base = format!("midi-daemon:{route_name}");
let is_default = decl.is_default();
let mut out_conns = HashMap::new();
let mut midi_fwds = HashMap::new();
let mut in_conns = Vec::new();
for port_name in &decl.outputs {
let (client_name, alsa_port) = if is_default {
(format!("{base}-out"), base.clone())
} else {
(
format!("{base}/{port_name}-out"),
format!("{base}/{port_name}"),
)
};
let midi_out =
MidiOutput::new(&client_name).context("Failed to create MIDI output")?;
let conn = midi_out
.create_virtual(&alsa_port)
.map_err(|e| anyhow::anyhow!("Failed to create virtual MIDI output '{alsa_port}': {e}"))?;
out_conns.insert(port_name.clone(), Arc::new(Mutex::new(conn)));
}
for port_name in &decl.inputs {
let fwd: Arc<Mutex<Option<mpsc::Sender<RouteEvent>>>> =
Arc::new(Mutex::new(Some(initial_tx.clone())));
let alsa_name = if is_default {
format!("{base}-in")
} else {
format!("{base}/{port_name}-in")
};
let midi_in = MidiInput::new(&alsa_name).context("Failed to create MIDI input")?;
let fwd_ref = Arc::clone(&fwd);
let port_name_owned = port_name.clone();
let in_conn = midi_in
.create_virtual(
&alsa_name,
move |_stamp, message, ()| {
let guard = fwd_ref.lock().unwrap_or_else(std::sync::PoisonError::into_inner);
if let Some(tx) = guard.as_ref()
&& tx.try_send(RouteEvent::Midi {
port: port_name_owned.clone(),
bytes: message.to_vec(),
}).is_err() {
warn!("MIDI message dropped: route event channel full or closed");
}
},
(),
)
.map_err(|e| anyhow::anyhow!("Failed to create virtual MIDI input '{alsa_name}': {e}"))?;
midi_fwds.insert(port_name.clone(), fwd);
in_conns.push(in_conn);
}
Ok(Rc::new(RoutePorts {
out_conns,
midi_fwds,
_in_conns: in_conns,
decl: decl.clone(),
}))
}
fn redirect_inputs(&self, new_tx: &mpsc::Sender<RouteEvent>) {
for fwd in self.midi_fwds.values() {
*fwd.lock().unwrap_or_else(std::sync::PoisonError::into_inner) = Some(new_tx.clone());
}
}
}
pub struct Route {
ports: Rc<RoutePorts>,
_timer: Arc<Timer>,
thread: Option<std::thread::JoinHandle<()>>,
osc_tx: mpsc::Sender<RouteEvent>,
pub connect_decl: ConnectDecl,
pub osc_receive_port: Option<u16>,
pub obs_connections: Vec<String>,
}
impl Route {
pub fn send_resync(&self) {
let _ = self.osc_tx.try_send(RouteEvent::ResyncState);
}
pub fn shutdown(mut self) -> Option<std::thread::JoinHandle<()>> {
let _ = self.osc_tx.try_send(RouteEvent::Shutdown);
self.thread.take()
}
pub fn make_osc_injector(
&self,
) -> impl Fn(SocketAddr, String, Vec<rosc::OscType>) + Send + 'static {
let tx = self.osc_tx.clone();
move |from: SocketAddr, address: String, args: Vec<rosc::OscType>| {
if tx.try_send(RouteEvent::Osc { from, address, args }).is_err() {
warn!("OSC message dropped: route event channel full or closed");
}
}
}
pub fn make_obs_injector(&self) -> impl Fn(&str, serde_json::Value) + Send + Sync + 'static {
let tx = self.osc_tx.clone();
move |connection: &str, event: serde_json::Value| {
if tx
.try_send(RouteEvent::ObsEvent { connection: connection.to_string(), event })
.is_err()
{
warn!("OBS event dropped: route event channel full or closed");
}
}
}
}
impl Route {
pub fn ports_rc(&self) -> Rc<RoutePorts> {
Rc::clone(&self.ports)
}
pub fn port_decl(&self) -> &PortDecl {
&self.ports.decl
}
#[allow(clippy::too_many_lines)]
pub fn start(
lua_path: &Path,
config: &Arc<Config>,
existing_ports: Option<Rc<RoutePorts>>,
obs_manager: &Arc<crate::obs::ObsManager>,
) -> Result<Self> {
let name = lua_path
.file_stem()
.and_then(|s| s.to_str())
.unwrap_or("unknown")
.to_string();
let script = std::fs::read_to_string(lua_path)
.with_context(|| format!("Failed to read {}", lua_path.display()))?;
let route_cfg = config.route_config(&name).cloned();
let lua_state_file = config.lua_state_dir().join(&name).join("state.json");
let (decl, raw_connect, osc_decl, obs_decl) =
extract_all_decls(&script, &name, route_cfg.as_ref(), &config.routes_dir)?;
let connect_decl = apply_connect_defaults(
raw_connect,
&decl.inputs,
&decl.outputs,
config.default_connect_input.as_deref(),
config.default_connect_output.as_deref(),
);
let (tx, rx) = mpsc::channel::<RouteEvent>(256);
let ports = match existing_ports {
Some(p) if p.decl == decl => {
p.redirect_inputs(&tx);
p
}
other => {
if other.is_some() {
warn!(
"Route '{}': port layout changed on reload — ALSA port IDs will change",
name
);
}
RoutePorts::create(&name, &decl, &tx)?
}
};
let OscDecl { receive_port: osc_receive_port, send_targets: osc_send_targets } = osc_decl;
let any_receive = osc_receive_port.is_some() || config.osc_receive_port.is_some();
let osc_sender = if !osc_send_targets.is_empty() {
match OscSender::new(osc_send_targets) {
Ok(s) => Some(s),
Err(e) => {
warn!("Route '{}': failed to create OSC sender: {}", name, e);
None
}
}
} else if let Some(addr) = config.osc_send_addr.as_deref().and_then(parse_socket_addr) {
let mut targets = HashMap::new();
targets.insert("default".to_string(), addr);
match OscSender::new(targets) {
Ok(s) => Some(s),
Err(e) => {
warn!("Route '{}': failed to create OSC sender from global config: {}", name, e);
None
}
}
} else if any_receive {
match OscSender::new(HashMap::new()) {
Ok(s) => Some(s),
Err(e) => {
warn!("Route '{}': failed to create OSC reply socket: {}", name, e);
None
}
}
} else {
None
};
let osc_tx = tx.clone();
let obs_senders = obs_manager.sender_map();
let timer = Arc::new(Timer::new(config.default_bpm, config.default_ppqn));
let _timer_thread = timer.spawn(tx.clone(), RouteEvent::Timer);
let out_conns_for_thread: HashMap<String, Arc<Mutex<MidiOutputConnection>>> = ports
.out_conns
.iter()
.map(|(k, v)| (k.clone(), Arc::clone(v)))
.collect();
let default_out = ports.decl.outputs.first().cloned().unwrap_or_default();
let timer_for_thread = Arc::clone(&timer);
let name_for_thread = name.clone();
let osc_heartbeat_interval = config.osc_heartbeat_interval;
let routes_dir = config.routes_dir.clone();
let thread = std::thread::spawn(move || {
if let Err(e) = run_lua_event_loop(
&name_for_thread,
&script,
rx,
out_conns_for_thread,
default_out,
&timer_for_thread,
RouteThreadArgs { route_cfg, osc_sender, osc_heartbeat_interval, lua_state_file, obs_senders, routes_dir },
) {
error!("Route '{}' event loop error: {}", name_for_thread, e);
}
});
info!(
"Started route '{}' — inputs: [{}], outputs: [{}]",
name,
ports.decl.inputs.join(", "),
ports.decl.outputs.join(", "),
);
Ok(Route {
ports,
_timer: timer,
thread: Some(thread),
osc_tx,
connect_decl,
osc_receive_port,
obs_connections: obs_decl.connections,
})
}
}
const LUA_STDLIB: &str = include_str!("lua/stdlib.lua");
fn setup_extract_lua(
lua: &Lua,
name: &str,
route_cfg: Option<&toml::Table>,
routes_dir: &Path,
) -> Result<()> {
lua.globals().set("send", lua.create_function(|_, _: LuaMultiValue| Ok(()))?)?;
lua.globals().set("send_osc", lua.create_function(|_, _: LuaMultiValue| Ok(()))?)?;
lua.globals().set("set_bpm", lua.create_function(|_, _: f64| Ok(()))?)?;
lua.globals().set("get_bpm", lua.create_function(|_, ()| -> LuaResult<f64> { Ok(120.0) })?)?;
lua.globals().set("set_ppqn", lua.create_function(|_, _: u32| Ok(()))?)?;
lua.globals().set("get_ppqn", lua.create_function(|_, ()| -> LuaResult<u32> { Ok(24) })?)?;
lua.globals().set("log", lua.create_function(|_, _: String| Ok(()))?)?;
lua.globals().set("save_state", lua.create_function(|_, _: LuaTable| Ok(()))?)?;
lua.globals().set("load_state", lua.create_function(|lua, ()| lua.create_table())?)?;
lua.globals().set("obs_call", lua.create_function(|_, _: LuaMultiValue| Ok(()))?)?;
lua.globals().set(
"obs_call_sync",
lua.create_function(|_, _: LuaMultiValue| -> LuaResult<(bool, LuaValue, LuaValue)> {
Ok((false, LuaValue::Nil, LuaValue::Nil))
})?,
)?;
lua.globals().set("ROUTE_NAME", name)?;
lua.globals().set("OSC_SEND_ENABLED", false)?;
lua.globals().set("ROUTES_DIR", routes_dir.to_string_lossy().into_owned())?;
let cfg_table = match route_cfg {
Some(tbl) => toml_table_to_lua(lua, tbl)
.map_err(|e| anyhow::anyhow!("Failed to convert config to Lua: {e}"))?,
None => lua.create_table()?,
};
lua.globals().set("config", cfg_table)?;
lua.load(LUA_STDLIB).set_name("stdlib").exec()
.map_err(|e| anyhow::anyhow!("Failed to load Lua stdlib: {e}"))?;
Ok(())
}
fn lua_val_to_patterns(val: LuaValue) -> Vec<String> {
match val {
LuaValue::String(s) => s.to_str().ok().map(|p| vec![p.to_string()]).unwrap_or_default(),
LuaValue::Table(t) => {
let mut pats = Vec::new();
for i in 1u32.. {
match t.get::<LuaValue>(i) {
Ok(LuaValue::String(s)) => {
if let Ok(p) = s.to_str() { pats.push(p.to_string()); }
}
_ => break,
}
}
pats
}
_ => vec![],
}
}
fn connect_from_lua_table(tbl: &LuaTable) -> ConnectDecl {
let mut decl = ConnectDecl::default();
let Ok(LuaValue::Table(connect_tbl)) = tbl.get::<LuaValue>("connect") else { return decl };
if let Ok(LuaValue::Table(t)) = connect_tbl.get::<LuaValue>("inputs") {
for (k, v) in t.pairs::<String, LuaValue>().flatten() {
let pats = lua_val_to_patterns(v);
if !pats.is_empty() { decl.inputs.insert(k, pats); }
}
}
if let Ok(LuaValue::Table(t)) = connect_tbl.get::<LuaValue>("outputs") {
for (k, v) in t.pairs::<String, LuaValue>().flatten() {
let pats = lua_val_to_patterns(v);
if !pats.is_empty() { decl.outputs.insert(k, pats); }
}
}
if let Ok(v) = connect_tbl.get::<LuaValue>("input") {
let pats = lua_val_to_patterns(v);
if !pats.is_empty() { decl.inputs.insert(String::new(), pats); }
}
if let Ok(v) = connect_tbl.get::<LuaValue>("output") {
let pats = lua_val_to_patterns(v);
if !pats.is_empty() { decl.outputs.insert(String::new(), pats); }
}
decl
}
fn connect_from_toml(mut decl: ConnectDecl, route_cfg: Option<&toml::Table>) -> ConnectDecl {
if let Some(cfg) = route_cfg {
for (key, val) in cfg {
let pats: Vec<String> = match val {
toml::Value::String(s) => vec![s.clone()],
toml::Value::Array(arr) => arr.iter()
.filter_map(|v| v.as_str().map(str::to_string))
.collect(),
_ => continue,
};
if pats.is_empty() { continue; }
if key == "connect_input" {
decl.inputs.entry(String::new()).or_insert_with(|| pats);
} else if key == "connect_output" {
decl.outputs.entry(String::new()).or_insert_with(|| pats);
} else if let Some(port) = key
.strip_prefix("connect_")
.and_then(|s| s.strip_suffix("-in"))
.filter(|s| !s.is_empty())
{
decl.inputs.entry(port.to_string()).or_insert_with(|| pats);
} else if let Some(port) = key
.strip_prefix("connect_")
.and_then(|s| s.strip_suffix("-out"))
.filter(|s| !s.is_empty())
{
decl.outputs.entry(port.to_string()).or_insert_with(|| pats);
}
}
}
decl
}
fn extract_all_decls(
script: &str,
name: &str,
route_cfg: Option<&toml::Table>,
routes_dir: &Path,
) -> Result<(PortDecl, ConnectDecl, OscDecl, ObsDecl)> {
let lua = Lua::new();
setup_extract_lua(&lua, name, route_cfg, routes_dir)?;
if let Err(e) = lua.load(script).set_name(name).exec() {
tracing::debug!(
"[{}] extract_all_decls: script error (will be reported by event loop): {}",
name, e
);
return Ok((
PortDecl::default(),
connect_from_toml(ConnectDecl::default(), route_cfg),
OscDecl::default(),
ObsDecl::default(),
));
}
if let Ok(Some(f)) = lua.globals().get::<Option<LuaFunction>>("init") {
match f.call::<LuaValue>(()) {
Ok(LuaValue::Table(ref tbl)) => {
let port_decl = parse_port_decl_from_lua(tbl);
let connect_decl = connect_from_toml(connect_from_lua_table(tbl), route_cfg);
let osc_decl = osc_from_lua_table(tbl);
let obs_decl = obs_from_lua_table(tbl);
return Ok((port_decl, connect_decl, osc_decl, obs_decl));
}
Ok(_) => warn!("[{}] init() did not return a table; using default ports", name),
Err(e) => warn!("[{}] init() error: {}; using default ports", name, e),
}
}
let port_decl = route_cfg
.and_then(parse_port_decl_from_toml)
.unwrap_or_default();
let connect_decl = connect_from_toml(ConnectDecl::default(), route_cfg);
Ok((port_decl, connect_decl, OscDecl::default(), ObsDecl::default()))
}
fn obs_from_lua_table(tbl: &LuaTable) -> ObsDecl {
let Ok(LuaValue::Table(obs_tbl)) = tbl.get::<LuaValue>("obs") else { return ObsDecl::default() };
let connections = match obs_tbl.get::<LuaValue>("connections") {
Ok(v) => lua_val_to_patterns(v),
_ => vec![],
};
ObsDecl { connections }
}
fn osc_from_lua_table(tbl: &LuaTable) -> OscDecl {
let mut decl = OscDecl::default();
let Ok(LuaValue::Table(osc_tbl)) = tbl.get::<LuaValue>("osc") else { return decl };
let port_num: Option<i64> = match osc_tbl.get::<LuaValue>("receive").unwrap_or(LuaValue::Nil) {
LuaValue::Integer(n) => Some(n),
#[allow(clippy::cast_possible_truncation)]
LuaValue::Number(f) if f.fract() == 0.0 => Some(f as i64),
LuaValue::Number(f) => {
warn!("OSC receive port must be an integer, got {}", f);
None
}
LuaValue::Nil => None,
other => {
warn!("OSC receive port must be an integer, got {}", other.type_name());
None
}
};
if let Some(n) = port_num {
if n > 0 && n <= 65535 {
decl.receive_port = Some(u16::try_from(n).expect("just checked n is in 1..=65535"));
} else {
warn!("OSC receive port {} is out of range (1–65535)", n);
}
}
if let Ok(LuaValue::Table(send_tbl)) = osc_tbl.get::<LuaValue>("send") {
for pair in send_tbl.pairs::<String, LuaValue>() {
if let Ok((target_name, LuaValue::String(addr_str))) = pair
&& let Ok(s) = addr_str.to_str() {
if let Some(addr) = parse_socket_addr(&s) {
decl.send_targets.insert(target_name, addr);
} else { warn!("OSC send target '{}': invalid address '{}'", target_name, s) }
}
}
}
decl
}
fn parse_socket_addr(s: &str) -> Option<SocketAddr> {
s.parse::<SocketAddr>().ok()
}
fn apply_connect_defaults(
mut raw: ConnectDecl,
port_inputs: &[String],
port_outputs: &[String],
global_input: Option<&str>,
global_output: Option<&str>,
) -> ConnectDecl {
let has_in = !raw.inputs.is_empty();
let has_out = !raw.outputs.is_empty();
let all_in = raw.inputs.remove("").or_else(|| if has_in { None } else { global_input.map(|s| vec![s.to_string()]) });
let all_out = raw.outputs.remove("").or_else(|| if has_out { None } else { global_output.map(|s| vec![s.to_string()]) });
for port in port_inputs {
if !raw.inputs.contains_key(port)
&& let Some(ref pats) = all_in {
raw.inputs.insert(port.clone(), pats.clone());
}
}
for port in port_outputs {
if !raw.outputs.contains_key(port)
&& let Some(ref pats) = all_out {
raw.outputs.insert(port.clone(), pats.clone());
}
}
raw
}
fn parse_port_decl_from_lua(tbl: &LuaTable) -> PortDecl {
fn extract_names(val: LuaValue) -> Vec<String> {
match val {
LuaValue::String(s) => vec![s.to_str().map(|b| b.to_string()).unwrap_or_default()],
LuaValue::Table(t) => {
let mut names = Vec::new();
for i in 1u32.. {
match t.get::<LuaValue>(i) {
Ok(LuaValue::String(s)) => {
names.push(s.to_str().map(|b| b.to_string()).unwrap_or_default());
}
_ => break,
}
}
names
}
_ => vec![],
}
}
let inputs = extract_names(tbl.get::<LuaValue>("inputs").unwrap_or(LuaValue::Nil));
let outputs = extract_names(tbl.get::<LuaValue>("outputs").unwrap_or(LuaValue::Nil));
if inputs.is_empty() || outputs.is_empty() {
return PortDecl::default();
}
PortDecl { inputs, outputs }
}
fn parse_port_decl_from_toml(cfg: &toml::Table) -> Option<PortDecl> {
let to_strings = |arr: &[toml::Value]| -> Vec<String> {
arr.iter()
.filter_map(|v| v.as_str().map(str::to_string))
.collect()
};
let inputs = to_strings(cfg.get("inputs")?.as_array()?);
let outputs = to_strings(cfg.get("outputs")?.as_array()?);
if inputs.is_empty() || outputs.is_empty() {
return None;
}
Some(PortDecl { inputs, outputs })
}
struct RouteThreadArgs {
route_cfg: Option<toml::Table>,
osc_sender: Option<OscSender>,
osc_heartbeat_interval: f64,
lua_state_file: std::path::PathBuf,
obs_senders: HashMap<String, mpsc::Sender<crate::obs::ObsRequest>>,
routes_dir: std::path::PathBuf,
}
#[allow(clippy::too_many_lines)]
fn run_lua_event_loop(
name: &str,
script: &str,
mut rx: mpsc::Receiver<RouteEvent>,
out_conns: HashMap<String, Arc<Mutex<MidiOutputConnection>>>,
default_out: String,
timer: &Arc<Timer>,
args: RouteThreadArgs,
) -> Result<()> {
let RouteThreadArgs { route_cfg, osc_sender, osc_heartbeat_interval, lua_state_file, obs_senders, routes_dir } = args;
let lua = Lua::new();
register_send(&lua, out_conns, default_out)?;
register_timer_fns(&lua, timer)?;
register_log(&lua, name)?;
register_obs_fns(&lua, obs_senders)?;
lua.globals().set("ROUTE_NAME", name)?;
lua.globals().set("OSC_SEND_ENABLED", osc_sender.is_some())?;
lua.globals().set("ROUTES_DIR", routes_dir.to_string_lossy().into_owned())?;
let subs_cache = register_send_osc(&lua, osc_sender)?;
{
let cfg_table = match route_cfg {
Some(ref tbl) => toml_table_to_lua(&lua, tbl)
.map_err(|e| anyhow::anyhow!("Failed to convert route config to Lua: {e}"))?,
None => lua.create_table()?,
};
lua.globals().set("config", cfg_table)?;
}
register_state_functions(&lua, &lua_state_file, name)?;
lua.load(LUA_STDLIB).set_name("stdlib").exec()
.map_err(|e| anyhow::anyhow!("Failed to load Lua stdlib: {e}"))?;
anyhow::Context::with_context(lua.load(script).set_name(name).exec(), || {
format!("Lua load error in '{name}'")
})?;
let mut osc_param_set: Option<crate::osc_params::OscParamSet> =
lua.globals()
.get::<Option<LuaFunction>>("init")
.ok()
.flatten()
.and_then(|f| f.call::<LuaValue>(()).ok())
.and_then(|v| if let LuaValue::Table(t) = v { Some(t) } else { None })
.and_then(|tbl| {
let prefix = format!("/{name}");
match crate::osc_params::from_init_table(&lua, &prefix, &tbl, osc_heartbeat_interval) {
Ok(ps) => ps,
Err(e) => {
warn!("[{}] Failed to build OscParamSet from init(): {}", name, e);
None
}
}
});
let on_midi_fn: Option<LuaFunction> = lua.globals().get("on_midi").ok();
let on_tick_fn: Option<LuaFunction> = lua.globals().get("on_tick").ok();
let on_osc_fn: Option<LuaFunction> = lua.globals().get("on_osc").ok();
let on_obs_event_fn: Option<LuaFunction> = lua.globals().get("on_obs_event").ok();
call_optional_hook(&lua, "on_startup", name);
while let Some(event) = rx.blocking_recv() {
match event {
RouteEvent::Midi { port, bytes } => {
let needs_parse = on_midi_fn.is_some() || osc_param_set.is_some();
if needs_parse {
match midi_bytes_to_lua(&lua, &bytes) {
Ok(msg) => {
let _ = msg.set("port", port.as_str());
if let Some(ref mut ps) = osc_param_set
&& let Err(e) = ps.dispatch_midi(&lua, &msg) {
warn!("[{}] midi param dispatch error: {}", name, e);
}
if let Some(ref on_midi) = on_midi_fn
&& let Err(e) = on_midi.call::<()>(msg) {
warn!("[{}] on_midi error: {}", name, e);
}
}
Err(e) => warn!("[{}] MIDI parse error: {}", name, e),
}
}
}
RouteEvent::Timer(TimerEvent::Tick { tick, bpm, ppqn }) => {
if let Some(ref mut ps) = osc_param_set {
if let Err(e) = ps.tick(&lua) {
warn!("[{}] osc_params tick error: {}", name, e);
}
*subs_cache.lock().unwrap_or_else(std::sync::PoisonError::into_inner) = ps.subscriber_addrs();
}
if let Some(ref on_tick) = on_tick_fn
&& let Err(e) = on_tick.call::<()>((tick, bpm, ppqn)) {
warn!("[{}] on_tick error: {}", name, e);
}
}
RouteEvent::Osc { from, address, args } => {
match osc_message_to_lua(&lua, &address, &args) {
Ok(msg) => {
let _ = msg.set("from", from.to_string());
if let Some(ref mut ps) = osc_param_set {
if let Err(e) = ps.dispatch(&lua, &msg) {
warn!("[{}] osc_params dispatch error: {}", name, e);
}
*subs_cache.lock().unwrap_or_else(std::sync::PoisonError::into_inner) = ps.subscriber_addrs();
}
if let Some(ref on_osc) = on_osc_fn {
if let Err(e) = on_osc.call::<()>(msg) {
warn!("[{}] on_osc error: {}", name, e);
}
} else if osc_param_set.is_none() {
warn!("[{}] OSC message '{}' received but no on_osc handler defined", name, address);
}
}
Err(e) => warn!("[{}] OSC message parse error: {}", name, e),
}
}
RouteEvent::ObsEvent { connection, event } => {
if let Some(ref on_obs_event) = on_obs_event_fn {
match lua.to_value(&event) {
Ok(ev_lua) => {
if let Err(e) = on_obs_event.call::<()>((connection.as_str(), ev_lua)) {
warn!("[{}] on_obs_event error: {}", name, e);
}
}
Err(e) => warn!("[{}] OBS event convert error: {}", name, e),
}
}
}
RouteEvent::ResyncState => {
if let Some(ref ps) = osc_param_set
&& let Err(e) = ps.resync(&lua) {
warn!("[{}] resync error: {}", name, e);
}
}
RouteEvent::Shutdown => break,
}
}
call_optional_hook(&lua, "on_shutdown", name);
Ok(())
}
fn register_send(
lua: &Lua,
out_conns: HashMap<String, Arc<Mutex<MidiOutputConnection>>>,
default_out: String,
) -> LuaResult<()> {
let send_fn = lua.create_function(move |_lua, args: LuaMultiValue| -> LuaResult<()> {
let (port_name, msg_table) = match args.len() {
1 => {
let Some(LuaValue::Table(msg)) = args.into_iter().next() else {
return Err(LuaError::RuntimeError(
"send: expected a message table".into(),
));
};
(default_out.clone(), msg)
}
2 => {
let mut iter = args.into_iter();
let Some(LuaValue::String(s)) = iter.next() else {
return Err(LuaError::RuntimeError(
"send: first argument must be a port name string".into(),
));
};
let port = s.to_str().map_err(LuaError::external)?.to_string();
let Some(LuaValue::Table(msg)) = iter.next() else {
return Err(LuaError::RuntimeError(
"send: second argument must be a message table".into(),
));
};
(port, msg)
}
n => {
return Err(LuaError::RuntimeError(format!(
"send: expected 1 or 2 arguments, got {n}"
)))
}
};
if let Some(conn) = out_conns.get(&port_name) { match lua_to_midi_bytes(&msg_table) {
Ok(bytes) => {
if let Err(e) = conn.lock().unwrap_or_else(std::sync::PoisonError::into_inner).send(&bytes) {
warn!("MIDI send error on port '{}': {}", port_name, e);
}
}
Err(e) => warn!("lua_to_midi_bytes error: {}", e),
} } else { warn!("send: unknown output port '{}'", port_name) }
Ok(())
})?;
lua.globals().set("send", send_fn)
}
fn register_timer_fns(lua: &Lua, timer: &Arc<Timer>) -> LuaResult<()> {
{
let t = Arc::clone(timer);
let f = lua.create_function(move |_, bpm: f64| {
t.set_bpm(bpm);
Ok(())
})?;
lua.globals().set("set_bpm", f)?;
}
{
let t = Arc::clone(timer);
let f = lua.create_function(move |_, ()| Ok(t.get_bpm()))?;
lua.globals().set("get_bpm", f)?;
}
{
let t = Arc::clone(timer);
let f = lua.create_function(move |_, ppqn: u32| {
t.set_ppqn(ppqn);
Ok(())
})?;
lua.globals().set("set_ppqn", f)?;
}
{
let t = Arc::clone(timer);
let f = lua.create_function(move |_, ()| Ok(t.get_ppqn()))?;
lua.globals().set("get_ppqn", f)?;
}
Ok(())
}
fn register_log(lua: &Lua, route_name: &str) -> LuaResult<()> {
let route_name = route_name.to_string();
let f = lua.create_function(move |_, msg: String| {
info!("[{}] {}", route_name, msg);
Ok(())
})?;
lua.globals().set("log", f)
}
fn register_send_osc(
lua: &Lua,
osc_sender: Option<OscSender>,
) -> LuaResult<Arc<Mutex<Vec<String>>>> {
let subs_cache: Arc<Mutex<Vec<String>>> = Arc::new(Mutex::new(Vec::new()));
let subs_cache_for_send = Arc::clone(&subs_cache);
let f = lua.create_function(move |_, args: LuaMultiValue| -> LuaResult<()> {
let Some(sender) = &osc_sender else {
warn!("send_osc: no OSC socket available");
return Ok(());
};
if args.is_empty() {
return Err(LuaError::RuntimeError(
"send_osc: expected at least an OSC address argument".into(),
));
}
let first = match &args[0] {
LuaValue::String(s) => s.to_str().map_err(LuaError::external)?.to_string(),
_ => return Err(LuaError::RuntimeError(
"send_osc: first argument must be a string".into(),
)),
};
if let Ok(dest) = first.parse::<SocketAddr>() {
let address = match args.get(1) {
Some(LuaValue::String(s)) => s.to_str().map_err(LuaError::external)?.to_string(),
_ => return Err(LuaError::RuntimeError(
"send_osc: OSC address (second argument) must be a string".into(),
)),
};
if !address.starts_with('/') {
return Err(LuaError::RuntimeError(format!(
"send_osc: OSC address must start with '/', got '{address}'"
)));
}
let osc_args = args.into_iter().skip(2)
.map(|v| lua_val_to_osc_type(&v))
.collect::<LuaResult<Vec<_>>>()?;
if let Err(e) = sender.send_to_addr(dest, address, osc_args) {
warn!("OSC send error sending to {}: {}", dest, e);
}
return Ok(());
}
let (target, address, arg_start) = if first.starts_with('/') {
if sender.targets.len() == 1 {
let t = sender.targets.keys().next().unwrap().clone();
(t, first, 1usize)
} else if sender.targets.is_empty() {
let subs: Vec<String> = subs_cache_for_send.lock().unwrap_or_else(std::sync::PoisonError::into_inner).clone();
if !subs.is_empty() {
let osc_args = args.into_iter().skip(1)
.map(|v| lua_val_to_osc_type(&v))
.collect::<LuaResult<Vec<_>>>()?;
for sub_addr in &subs {
if let Ok(dest) = sub_addr.parse::<SocketAddr>()
&& let Err(e) = sender.send_to_addr(dest, first.clone(), osc_args.clone()) {
warn!("OSC send error sending to {}: {}", dest, e);
}
}
}
return Ok(());
} else {
return Err(LuaError::RuntimeError(
"send_osc: multiple targets configured — specify target name as first argument".into(),
));
}
} else {
if !sender.targets.contains_key(&first) {
return Err(LuaError::RuntimeError(format!(
"send_osc: unknown target '{first}'"
)));
}
let address = match args.get(1) {
Some(LuaValue::String(s)) => s.to_str().map_err(LuaError::external)?.to_string(),
_ => return Err(LuaError::RuntimeError(
"send_osc: OSC address (second argument) must be a string".into(),
)),
};
if !address.starts_with('/') {
return Err(LuaError::RuntimeError(format!(
"send_osc: OSC address must start with '/', got '{address}'"
)));
}
(first, address, 2usize)
};
let osc_args = args.into_iter()
.skip(arg_start)
.map(|v| lua_val_to_osc_type(&v))
.collect::<LuaResult<Vec<_>>>()?;
if let Err(e) = sender.send(&target, address, osc_args) {
warn!("OSC send error sending to '{}': {}", target, e);
}
Ok(())
})?;
lua.globals().set("send_osc", f)?;
Ok(subs_cache)
}
fn register_obs_fns(
lua: &Lua,
obs_senders: HashMap<String, mpsc::Sender<crate::obs::ObsRequest>>,
) -> LuaResult<()> {
let senders = Arc::new(obs_senders);
{
let senders = Arc::clone(&senders);
let f = lua.create_function(
move |lua, (conn, request, args): (String, String, LuaValue)| -> LuaResult<()> {
let Some(tx) = senders.get(&conn) else {
warn!("obs_call: unknown OBS connection '{}'", conn);
return Ok(());
};
let args_json: serde_json::Value = lua.from_value(args)?;
let req = crate::obs::ObsRequest { request, args: args_json, reply: None };
if tx.try_send(req).is_err() {
warn!("obs_call: request dropped — connection '{}' queue full or closed", conn);
}
Ok(())
},
)?;
lua.globals().set("obs_call", f)?;
}
{
let f = lua.create_function(
move |lua, (conn, request, args, timeout_ms): (String, String, LuaValue, u64)| -> LuaResult<(bool, LuaValue, LuaValue)> {
let Some(tx) = senders.get(&conn) else {
let err = lua.create_string(format!("unknown OBS connection '{conn}'"))?;
return Ok((false, LuaValue::Nil, LuaValue::String(err)));
};
let args_json: serde_json::Value = lua.from_value(args)?;
let (reply_tx, reply_rx) = std::sync::mpsc::channel();
let req = crate::obs::ObsRequest { request, args: args_json, reply: Some(reply_tx) };
if tx.try_send(req).is_err() {
let err = lua.create_string("request queue full or closed")?;
return Ok((false, LuaValue::Nil, LuaValue::String(err)));
}
if let Ok(reply) = reply_rx.recv_timeout(std::time::Duration::from_millis(timeout_ms)) {
let result_lua = lua.to_value(&reply.result)?;
let err_lua = match reply.error {
Some(e) => LuaValue::String(lua.create_string(&e)?),
None => LuaValue::Nil,
};
Ok((reply.ok, result_lua, err_lua))
} else {
let err = lua.create_string("obs_call_sync: timed out")?;
Ok((false, LuaValue::Nil, LuaValue::String(err)))
}
},
)?;
lua.globals().set("obs_call_sync", f)?;
}
Ok(())
}
fn register_state_functions(lua: &Lua, path: &Path, route_name: &str) -> LuaResult<()> {
{
let path = path.to_path_buf();
let route_name = route_name.to_string();
let f = lua.create_function(move |lua, table: LuaTable| -> LuaResult<()> {
if let Err(e) = crate::lua_api::save_json_state(lua, &path, &table) {
warn!("[{}] save_state error: {}", route_name, e);
}
Ok(())
})?;
lua.globals().set("save_state", f)?;
}
{
let path = path.to_path_buf();
let route_name = route_name.to_string();
let f = lua.create_function(move |lua, ()| -> LuaResult<LuaTable> {
Ok(crate::lua_api::load_json_state(lua, &path).unwrap_or_else(|e| {
warn!("[{}] load_state error: {}", route_name, e);
lua.create_table().expect("create empty Lua state table")
}))
})?;
lua.globals().set("load_state", f)?;
}
Ok(())
}
fn call_optional_hook(lua: &Lua, hook_name: &str, route_name: &str) {
if let Ok(Some(f)) = lua.globals().get::<Option<LuaFunction>>(hook_name)
&& let Err(e) = f.call::<()>(()) {
warn!("[{}] {} error: {}", route_name, hook_name, e);
}
}
#[cfg(test)]
mod tests {
use super::*;
use mlua::Lua;
#[test]
fn port_decl_default_has_single_default_input_and_output() {
let d = PortDecl::default();
assert_eq!(d.inputs, vec!["default"]);
assert_eq!(d.outputs, vec!["default"]);
}
#[test]
fn port_decl_is_default_true_for_default() {
assert!(PortDecl::default().is_default());
}
#[test]
fn port_decl_is_default_false_for_custom_input_name() {
let d = PortDecl {
inputs: vec!["keyboard".to_string()],
outputs: vec!["default".to_string()],
};
assert!(!d.is_default());
}
#[test]
fn port_decl_is_default_false_for_custom_output_name() {
let d = PortDecl {
inputs: vec!["default".to_string()],
outputs: vec!["synth".to_string()],
};
assert!(!d.is_default());
}
#[test]
fn port_decl_is_default_false_for_multiple_inputs() {
let d = PortDecl {
inputs: vec!["default".to_string(), "extra".to_string()],
outputs: vec!["default".to_string()],
};
assert!(!d.is_default());
}
#[test]
fn port_decl_is_default_false_for_multiple_outputs() {
let d = PortDecl {
inputs: vec!["default".to_string()],
outputs: vec!["default".to_string(), "extra".to_string()],
};
assert!(!d.is_default());
}
#[test]
fn port_decl_is_default_false_for_empty_inputs() {
let d = PortDecl {
inputs: vec![],
outputs: vec!["default".to_string()],
};
assert!(!d.is_default());
}
#[test]
fn port_decl_equality_same_ports() {
let a = PortDecl {
inputs: vec!["kbd".to_string()],
outputs: vec!["synth".to_string()],
};
let b = a.clone();
assert_eq!(a, b);
}
#[test]
fn port_decl_inequality_different_input_names() {
let a = PortDecl {
inputs: vec!["kbd".to_string()],
outputs: vec!["synth".to_string()],
};
let b = PortDecl {
inputs: vec!["pad".to_string()],
outputs: vec!["synth".to_string()],
};
assert_ne!(a, b);
}
#[test]
fn port_decl_inequality_different_output_names() {
let a = PortDecl {
inputs: vec!["kbd".to_string()],
outputs: vec!["synth".to_string()],
};
let b = PortDecl {
inputs: vec!["kbd".to_string()],
outputs: vec!["drums".to_string()],
};
assert_ne!(a, b);
}
#[test]
fn port_decl_inequality_input_order_matters() {
let a = PortDecl {
inputs: vec!["x".to_string(), "y".to_string()],
outputs: vec!["z".to_string()],
};
let b = PortDecl {
inputs: vec!["y".to_string(), "x".to_string()],
outputs: vec!["z".to_string()],
};
assert_ne!(a, b);
}
#[test]
fn port_decl_inequality_output_order_matters() {
let a = PortDecl {
inputs: vec!["x".to_string()],
outputs: vec!["a".to_string(), "b".to_string()],
};
let b = PortDecl {
inputs: vec!["x".to_string()],
outputs: vec!["b".to_string(), "a".to_string()],
};
assert_ne!(a, b);
}
fn toml_section(s: &str) -> toml::Table {
toml::from_str(s).unwrap()
}
#[test]
fn toml_single_input_and_output() {
let tbl = toml_section("inputs = [\"kbd\"]\noutputs = [\"synth\"]");
let decl = parse_port_decl_from_toml(&tbl).unwrap();
assert_eq!(decl.inputs, vec!["kbd"]);
assert_eq!(decl.outputs, vec!["synth"]);
}
#[test]
fn toml_multiple_inputs_and_outputs() {
let tbl = toml_section("inputs = [\"kbd\", \"pad\"]\noutputs = [\"synth\", \"drums\"]");
let decl = parse_port_decl_from_toml(&tbl).unwrap();
assert_eq!(decl.inputs, vec!["kbd", "pad"]);
assert_eq!(decl.outputs, vec!["synth", "drums"]);
}
#[test]
fn toml_missing_inputs_returns_none() {
let tbl = toml_section("outputs = [\"synth\"]");
assert!(parse_port_decl_from_toml(&tbl).is_none());
}
#[test]
fn toml_missing_outputs_returns_none() {
let tbl = toml_section("inputs = [\"kbd\"]");
assert!(parse_port_decl_from_toml(&tbl).is_none());
}
#[test]
fn toml_empty_inputs_array_returns_none() {
let tbl = toml_section("inputs = []\noutputs = [\"synth\"]");
assert!(parse_port_decl_from_toml(&tbl).is_none());
}
#[test]
fn toml_empty_outputs_array_returns_none() {
let tbl = toml_section("inputs = [\"kbd\"]\noutputs = []");
assert!(parse_port_decl_from_toml(&tbl).is_none());
}
#[test]
fn toml_empty_table_returns_none() {
let tbl = toml_section("");
assert!(parse_port_decl_from_toml(&tbl).is_none());
}
#[test]
fn toml_non_string_values_only_returns_none() {
let tbl = toml_section("inputs = [1, 2]\noutputs = [\"synth\"]");
assert!(parse_port_decl_from_toml(&tbl).is_none());
}
#[test]
fn toml_mixed_string_and_non_string_keeps_strings() {
let tbl = toml_section("inputs = [\"kbd\", 42]\noutputs = [\"synth\"]");
let decl = parse_port_decl_from_toml(&tbl).unwrap();
assert_eq!(decl.inputs, vec!["kbd"]);
}
#[test]
fn toml_preserves_port_order() {
let tbl = toml_section("inputs = [\"z\", \"a\", \"m\"]\noutputs = [\"out\"]");
let decl = parse_port_decl_from_toml(&tbl).unwrap();
assert_eq!(decl.inputs, vec!["z", "a", "m"]);
}
#[test]
fn toml_extra_keys_are_ignored() {
let tbl = toml_section(
"inputs = [\"kbd\"]\noutputs = [\"synth\"]\nbpm = 120\nchannel = 1",
);
let decl = parse_port_decl_from_toml(&tbl).unwrap();
assert_eq!(decl.inputs, vec!["kbd"]);
assert_eq!(decl.outputs, vec!["synth"]);
}
fn lua_array(lua: &Lua, items: &[&str]) -> LuaTable {
let t = lua.create_table().unwrap();
for (i, s) in items.iter().enumerate() {
t.set(i + 1, *s).unwrap();
}
t
}
fn lua_decl_table(lua: &Lua, inputs: &[&str], outputs: &[&str]) -> LuaTable {
let tbl = lua.create_table().unwrap();
tbl.set("inputs", lua_array(lua, inputs)).unwrap();
tbl.set("outputs", lua_array(lua, outputs)).unwrap();
tbl
}
#[test]
fn lua_single_input_and_output() {
let lua = Lua::new();
let tbl = lua_decl_table(&lua, &["kbd"], &["synth"]);
let decl = parse_port_decl_from_lua(&tbl);
assert_eq!(decl.inputs, vec!["kbd"]);
assert_eq!(decl.outputs, vec!["synth"]);
}
#[test]
fn lua_multiple_inputs_and_outputs() {
let lua = Lua::new();
let tbl = lua_decl_table(&lua, &["kbd", "pad"], &["synth", "drums"]);
let decl = parse_port_decl_from_lua(&tbl);
assert_eq!(decl.inputs, vec!["kbd", "pad"]);
assert_eq!(decl.outputs, vec!["synth", "drums"]);
}
#[test]
fn lua_string_shorthand_for_inputs() {
let lua = Lua::new();
let tbl = lua.create_table().unwrap();
tbl.set("inputs", "kbd").unwrap();
tbl.set("outputs", lua_array(&lua, &["synth"])).unwrap();
let decl = parse_port_decl_from_lua(&tbl);
assert_eq!(decl.inputs, vec!["kbd"]);
}
#[test]
fn lua_string_shorthand_for_outputs() {
let lua = Lua::new();
let tbl = lua.create_table().unwrap();
tbl.set("inputs", lua_array(&lua, &["kbd"])).unwrap();
tbl.set("outputs", "synth").unwrap();
let decl = parse_port_decl_from_lua(&tbl);
assert_eq!(decl.outputs, vec!["synth"]);
}
#[test]
fn lua_string_shorthand_for_both() {
let lua = Lua::new();
let tbl = lua.create_table().unwrap();
tbl.set("inputs", "kbd").unwrap();
tbl.set("outputs", "synth").unwrap();
let decl = parse_port_decl_from_lua(&tbl);
assert_eq!(decl.inputs, vec!["kbd"]);
assert_eq!(decl.outputs, vec!["synth"]);
}
#[test]
fn lua_missing_inputs_returns_default() {
let lua = Lua::new();
let tbl = lua.create_table().unwrap();
tbl.set("outputs", lua_array(&lua, &["synth"])).unwrap();
let decl = parse_port_decl_from_lua(&tbl);
assert!(decl.is_default());
}
#[test]
fn lua_missing_outputs_returns_default() {
let lua = Lua::new();
let tbl = lua.create_table().unwrap();
tbl.set("inputs", lua_array(&lua, &["kbd"])).unwrap();
let decl = parse_port_decl_from_lua(&tbl);
assert!(decl.is_default());
}
#[test]
fn lua_empty_table_returns_default() {
let lua = Lua::new();
let tbl = lua.create_table().unwrap();
let decl = parse_port_decl_from_lua(&tbl);
assert!(decl.is_default());
}
#[test]
fn lua_empty_inputs_array_returns_default() {
let lua = Lua::new();
let tbl = lua.create_table().unwrap();
tbl.set("inputs", lua.create_table().unwrap()).unwrap();
tbl.set("outputs", lua_array(&lua, &["synth"])).unwrap();
let decl = parse_port_decl_from_lua(&tbl);
assert!(decl.is_default());
}
#[test]
fn lua_empty_outputs_array_returns_default() {
let lua = Lua::new();
let tbl = lua.create_table().unwrap();
tbl.set("inputs", lua_array(&lua, &["kbd"])).unwrap();
tbl.set("outputs", lua.create_table().unwrap()).unwrap();
let decl = parse_port_decl_from_lua(&tbl);
assert!(decl.is_default());
}
#[test]
fn lua_preserves_port_order() {
let lua = Lua::new();
let tbl = lua_decl_table(&lua, &["z", "a", "m"], &["out"]);
let decl = parse_port_decl_from_lua(&tbl);
assert_eq!(decl.inputs, vec!["z", "a", "m"]);
}
#[test]
fn lua_non_string_in_array_stops_iteration() {
let lua = Lua::new();
let inp = lua.create_table().unwrap();
inp.set(1, "kbd").unwrap();
inp.set(2, 42i64).unwrap();
inp.set(3, "pad").unwrap(); let tbl = lua.create_table().unwrap();
tbl.set("inputs", inp).unwrap();
tbl.set("outputs", lua_array(&lua, &["synth"])).unwrap();
let decl = parse_port_decl_from_lua(&tbl);
assert_eq!(decl.inputs, vec!["kbd"]);
}
#[test]
fn lua_nil_value_stops_iteration() {
let lua = Lua::new();
let inp = lua.create_table().unwrap();
inp.set(1, "kbd").unwrap();
inp.set(3, "pad").unwrap();
let tbl = lua.create_table().unwrap();
tbl.set("inputs", inp).unwrap();
tbl.set("outputs", lua_array(&lua, &["synth"])).unwrap();
let decl = parse_port_decl_from_lua(&tbl);
assert_eq!(decl.inputs, vec!["kbd"]);
}
fn extract(script: &str) -> PortDecl {
extract_all_decls(script, "test", None, Path::new("/test/routes")).unwrap().0
}
fn extract_with_cfg(script: &str, cfg_toml: &str) -> PortDecl {
let tbl: toml::Table = toml::from_str(cfg_toml).unwrap();
extract_all_decls(script, "test", Some(&tbl), Path::new("/test/routes")).unwrap().0
}
#[test]
fn extract_no_init_no_config_returns_default() {
assert!(extract("-- no init").is_default());
}
#[test]
fn extract_init_returns_named_ports() {
let decl = extract(r#"
function init()
return { inputs = {"kbd", "pad"}, outputs = {"synth", "drums"} }
end
"#);
assert_eq!(decl.inputs, vec!["kbd", "pad"]);
assert_eq!(decl.outputs, vec!["synth", "drums"]);
}
#[test]
fn extract_init_single_port_each() {
let decl = extract(r#"
function init()
return { inputs = {"kbd"}, outputs = {"synth"} }
end
"#);
assert_eq!(decl.inputs, vec!["kbd"]);
assert_eq!(decl.outputs, vec!["synth"]);
}
#[test]
fn extract_init_string_shorthand_for_both() {
let decl = extract(r#"
function init()
return { inputs = "kbd", outputs = "synth" }
end
"#);
assert_eq!(decl.inputs, vec!["kbd"]);
assert_eq!(decl.outputs, vec!["synth"]);
}
#[test]
fn extract_init_string_shorthand_for_inputs_only() {
let decl = extract(r#"
function init()
return { inputs = "kbd", outputs = {"synth", "drums"} }
end
"#);
assert_eq!(decl.inputs, vec!["kbd"]);
assert_eq!(decl.outputs, vec!["synth", "drums"]);
}
#[test]
fn extract_init_many_ports() {
let decl = extract(r#"
function init()
return {
inputs = {"in1", "in2", "in3", "in4"},
outputs = {"out1", "out2", "out3"},
}
end
"#);
assert_eq!(decl.inputs, vec!["in1", "in2", "in3", "in4"]);
assert_eq!(decl.outputs, vec!["out1", "out2", "out3"]);
}
#[test]
fn extract_init_overrides_config_toml() {
let decl = extract_with_cfg(
r#"
function init()
return { inputs = {"lua-in"}, outputs = {"lua-out"} }
end
"#,
"inputs = [\"toml-in\"]\noutputs = [\"toml-out\"]",
);
assert_eq!(decl.inputs, vec!["lua-in"]);
assert_eq!(decl.outputs, vec!["lua-out"]);
}
#[test]
fn extract_config_toml_used_when_no_init() {
let decl = extract_with_cfg(
"-- no init",
"inputs = [\"kbd\"]\noutputs = [\"synth\"]",
);
assert_eq!(decl.inputs, vec!["kbd"]);
assert_eq!(decl.outputs, vec!["synth"]);
}
#[test]
fn extract_init_returning_nil_falls_back_to_toml() {
let decl = extract_with_cfg(
r"
function init()
return nil
end
",
"inputs = [\"kbd\"]\noutputs = [\"synth\"]",
);
assert_eq!(decl.inputs, vec!["kbd"]);
assert_eq!(decl.outputs, vec!["synth"]);
}
#[test]
fn extract_init_returning_non_table_falls_back_to_toml() {
let decl = extract_with_cfg(
r#"
function init()
return "not a table"
end
"#,
"inputs = [\"kbd\"]\noutputs = [\"synth\"]",
);
assert_eq!(decl.inputs, vec!["kbd"]);
assert_eq!(decl.outputs, vec!["synth"]);
}
#[test]
fn extract_init_returning_non_table_falls_back_to_default_when_no_toml() {
let decl = extract(r"
function init()
return 42
end
");
assert!(decl.is_default());
}
#[test]
fn extract_init_error_falls_back_to_toml() {
let decl = extract_with_cfg(
r#"
function init()
error("something went wrong")
end
"#,
"inputs = [\"kbd\"]\noutputs = [\"synth\"]",
);
assert_eq!(decl.inputs, vec!["kbd"]);
assert_eq!(decl.outputs, vec!["synth"]);
}
#[test]
fn extract_init_error_falls_back_to_default_when_no_toml() {
let decl = extract(r#"
function init()
error("oops")
end
"#);
assert!(decl.is_default());
}
#[test]
fn extract_script_syntax_error_returns_default() {
let decl = extract("this is ][ not valid lua");
assert!(decl.is_default());
}
#[test]
fn extract_init_empty_inputs_returns_default_not_toml() {
let decl = extract_with_cfg(
r#"
function init()
return { inputs = {}, outputs = {"synth"} }
end
"#,
"inputs = [\"kbd\"]\noutputs = [\"synth\"]",
);
assert!(decl.is_default());
}
#[test]
fn extract_init_empty_outputs_returns_default_not_toml() {
let decl = extract_with_cfg(
r#"
function init()
return { inputs = {"kbd"}, outputs = {} }
end
"#,
"inputs = [\"kbd\"]\noutputs = [\"synth\"]",
);
assert!(decl.is_default());
}
#[test]
fn extract_config_toml_empty_inputs_returns_default() {
let decl = extract_with_cfg("-- no init", "inputs = []\noutputs = [\"synth\"]");
assert!(decl.is_default());
}
#[test]
fn extract_config_toml_empty_outputs_returns_default() {
let decl = extract_with_cfg("-- no init", "inputs = [\"kbd\"]\noutputs = []");
assert!(decl.is_default());
}
#[test]
fn extract_backward_compat_script_without_init() {
let decl = extract(r"
function on_midi(msg)
send(msg)
end
function on_tick(tick, bpm, ppqn)
end
");
assert!(decl.is_default());
}
#[test]
fn extract_script_can_use_config_global_in_init() {
let decl = extract_with_cfg(
r#"
function init()
local n = config.label or "fallback"
return { inputs = {n}, outputs = {"out"} }
end
"#,
"label = \"my-input\"\ninputs = [\"toml-in\"]\noutputs = [\"toml-out\"]",
);
assert_eq!(decl.inputs, vec!["my-input"]);
assert_eq!(decl.outputs, vec!["out"]);
}
#[test]
fn extract_script_can_call_log_in_init() {
let decl = extract(r#"
function init()
log("setting up ports")
return { inputs = {"kbd"}, outputs = {"synth"} }
end
"#);
assert_eq!(decl.inputs, vec!["kbd"]);
assert_eq!(decl.outputs, vec!["synth"]);
}
#[test]
fn extract_top_level_code_runs_without_real_connections() {
let decl = extract(r#"
log("top-level init")
local x = get_bpm()
function init()
return { inputs = {"in"}, outputs = {"out"} }
end
"#);
assert_eq!(decl.inputs, vec!["in"]);
assert_eq!(decl.outputs, vec!["out"]);
}
fn connect(script: &str) -> ConnectDecl {
extract_all_decls(script, "test", None, Path::new("/test/routes")).unwrap().1
}
fn connect_with_cfg(script: &str, cfg_toml: &str) -> ConnectDecl {
let tbl: toml::Table = toml::from_str(cfg_toml).unwrap();
extract_all_decls(script, "test", Some(&tbl), Path::new("/test/routes")).unwrap().1
}
#[test]
fn connect_no_init_no_toml_returns_empty() {
let c = connect("-- no init");
assert!(c.inputs.is_empty());
assert!(c.outputs.is_empty());
}
fn pats(strs: &[&str]) -> Vec<String> {
strs.iter().map(std::string::ToString::to_string).collect()
}
#[test]
fn connect_lua_singular_input_stored_under_sentinel() {
let c = connect(r#"
function init()
return { inputs = {"kbd"}, outputs = {"synth"},
connect = { input = ".*KeyLab.*" } }
end
"#);
assert_eq!(c.inputs.get(""), Some(&pats(&[".*KeyLab.*"])));
}
#[test]
fn connect_lua_singular_output_stored_under_sentinel() {
let c = connect(r#"
function init()
return { inputs = {"kbd"}, outputs = {"synth"},
connect = { output = ".*Surge.*" } }
end
"#);
assert_eq!(c.outputs.get(""), Some(&pats(&[".*Surge.*"])));
}
#[test]
fn connect_lua_singular_input_array_stored_under_sentinel() {
let c = connect(r#"
function init()
return { inputs = {"kbd"}, outputs = {"synth"},
connect = { input = {".*KeyLab.*", ".*A-PRO.*"} } }
end
"#);
assert_eq!(c.inputs.get(""), Some(&pats(&[".*KeyLab.*", ".*A-PRO.*"])));
}
#[test]
fn connect_lua_per_port_inputs_stored_by_name() {
let c = connect(r#"
function init()
return { inputs = {"kbd", "pad"}, outputs = {"synth"},
connect = { inputs = { kbd = ".*KORG.*", pad = ".*Alesis.*" } } }
end
"#);
assert_eq!(c.inputs.get("kbd"), Some(&pats(&[".*KORG.*"])));
assert_eq!(c.inputs.get("pad"), Some(&pats(&[".*Alesis.*"])));
assert!(!c.inputs.contains_key(""));
}
#[test]
fn connect_lua_per_port_inputs_array_stored_by_name() {
let c = connect(r#"
function init()
return { inputs = {"kbd"}, outputs = {"synth"},
connect = { inputs = { kbd = {".*KORG.*", ".*KeyLab.*"} } } }
end
"#);
assert_eq!(c.inputs.get("kbd"), Some(&pats(&[".*KORG.*", ".*KeyLab.*"])));
}
#[test]
fn connect_lua_per_port_outputs_stored_by_name() {
let c = connect(r#"
function init()
return { inputs = {"kbd"}, outputs = {"synth", "drums"},
connect = { outputs = { synth = ".*Surge.*", drums = ".*DrumKit.*" } } }
end
"#);
assert_eq!(c.outputs.get("synth"), Some(&pats(&[".*Surge.*"])));
assert_eq!(c.outputs.get("drums"), Some(&pats(&[".*DrumKit.*"])));
}
#[test]
fn connect_lua_per_port_and_singular_both_stored() {
let c = connect(r#"
function init()
return { inputs = {"kbd", "pad"}, outputs = {"synth"},
connect = {
inputs = { kbd = ".*KORG.*" },
input = ".*Fallback.*",
} }
end
"#);
assert_eq!(c.inputs.get("kbd"), Some(&pats(&[".*KORG.*"])));
assert_eq!(c.inputs.get(""), Some(&pats(&[".*Fallback.*"])));
}
#[test]
fn connect_lua_no_connect_key_returns_empty() {
let c = connect(r#"
function init()
return { inputs = {"kbd"}, outputs = {"synth"} }
end
"#);
assert!(c.inputs.is_empty());
assert!(c.outputs.is_empty());
}
#[test]
fn connect_toml_connect_input_stored_under_sentinel() {
let c = connect_with_cfg("-- no init", "connect_input = \".*KeyLab.*\"");
assert_eq!(c.inputs.get(""), Some(&pats(&[".*KeyLab.*"])));
}
#[test]
fn connect_toml_connect_input_array_stored_under_sentinel() {
let c = connect_with_cfg("-- no init", "connect_input = [\".*KeyLab.*\", \".*A-PRO.*\"]");
assert_eq!(c.inputs.get(""), Some(&pats(&[".*KeyLab.*", ".*A-PRO.*"])));
}
#[test]
fn connect_toml_connect_output_stored_under_sentinel() {
let c = connect_with_cfg("-- no init", "connect_output = \".*Surge.*\"");
assert_eq!(c.outputs.get(""), Some(&pats(&[".*Surge.*"])));
}
#[test]
fn connect_toml_per_port_input_stored_by_name() {
let c = connect_with_cfg("-- no init", "\"connect_keyboard-in\" = \".*A-PRO.*\"");
assert_eq!(c.inputs.get("keyboard"), Some(&pats(&[".*A-PRO.*"])));
assert!(!c.inputs.contains_key(""));
}
#[test]
fn connect_toml_per_port_input_array_stored_by_name() {
let c = connect_with_cfg("-- no init", "\"connect_keyboard-in\" = [\".*A-PRO.*\", \".*KeyLab.*\"]");
assert_eq!(c.inputs.get("keyboard"), Some(&pats(&[".*A-PRO.*", ".*KeyLab.*"])));
}
#[test]
fn connect_toml_per_port_output_stored_by_name() {
let c = connect_with_cfg("-- no init", "\"connect_synth-out\" = \".*Surge.*\"");
assert_eq!(c.outputs.get("synth"), Some(&pats(&[".*Surge.*"])));
assert!(!c.outputs.contains_key(""));
}
#[test]
fn connect_toml_per_port_multiple_inputs() {
let c = connect_with_cfg(
"-- no init",
"\"connect_keyboard-in\" = \".*A-PRO.*\"\n\"connect_metronome-in\" = \".*metronome-out.*\"",
);
assert_eq!(c.inputs.get("keyboard"), Some(&pats(&[".*A-PRO.*"])));
assert_eq!(c.inputs.get("metronome"), Some(&pats(&[".*metronome-out.*"])));
}
#[test]
fn connect_toml_per_port_input_not_overridden_by_sentinel() {
let c = connect_with_cfg(
"-- no init",
"connect_input = \".*Fallback.*\"\n\"connect_keyboard-in\" = \".*A-PRO.*\"",
);
assert_eq!(c.inputs.get("keyboard"), Some(&pats(&[".*A-PRO.*"])));
assert_eq!(c.inputs.get(""), Some(&pats(&[".*Fallback.*"])));
}
#[test]
fn connect_lua_per_port_overrides_toml_per_port() {
let c = connect_with_cfg(
r#"
function init()
return { inputs = {"keyboard"}, outputs = {"pan"},
connect = { inputs = { keyboard = ".*LuaDevice.*" } } }
end
"#,
"\"connect_keyboard-in\" = \".*TomlDevice.*\"",
);
assert_eq!(c.inputs.get("keyboard"), Some(&pats(&[".*LuaDevice.*"])));
}
#[test]
fn connect_lua_singular_overrides_toml_sentinel() {
let c = connect_with_cfg(
r#"
function init()
return { inputs = {"kbd"}, outputs = {"synth"},
connect = { input = ".*LuaPattern.*" } }
end
"#,
"connect_input = \".*TomlPattern.*\"",
);
assert_eq!(c.inputs.get(""), Some(&pats(&[".*LuaPattern.*"])));
}
#[test]
fn connect_toml_used_when_no_lua_connect() {
let c = connect_with_cfg(
r#"
function init()
return { inputs = {"kbd"}, outputs = {"synth"} }
end
"#,
"connect_input = \".*TomlPattern.*\"\nconnect_output = \".*TomlOut.*\"",
);
assert_eq!(c.inputs.get(""), Some(&pats(&[".*TomlPattern.*"])));
assert_eq!(c.outputs.get(""), Some(&pats(&[".*TomlOut.*"])));
}
#[test]
fn connect_script_error_returns_empty_falls_back_to_toml() {
let c = connect_with_cfg(
"this is ][ not valid lua",
"connect_input = \".*Fallback.*\"",
);
assert_eq!(c.inputs.get(""), Some(&pats(&[".*Fallback.*"])));
}
fn ports(names: &[&str]) -> Vec<String> {
names.iter().map(std::string::ToString::to_string).collect()
}
fn decl_with_sentinel(sentinel: &str) -> ConnectDecl {
let mut d = ConnectDecl::default();
d.inputs.insert(String::new(), vec![sentinel.to_string()]);
d
}
#[test]
fn defaults_empty_decl_no_globals_stays_empty() {
let result = apply_connect_defaults(
ConnectDecl::default(), &ports(&["default"]), &ports(&["default"]), None, None,
);
assert!(result.inputs.is_empty());
assert!(result.outputs.is_empty());
}
#[test]
fn defaults_global_fills_all_ports() {
let result = apply_connect_defaults(
ConnectDecl::default(),
&ports(&["kbd", "pad"]),
&ports(&["synth"]),
Some(".*MyController.*"),
Some(".*MySynth.*"),
);
assert_eq!(result.inputs.get("kbd"), Some(&pats(&[".*MyController.*"])));
assert_eq!(result.inputs.get("pad"), Some(&pats(&[".*MyController.*"])));
assert_eq!(result.outputs.get("synth"), Some(&pats(&[".*MySynth.*"])));
}
#[test]
fn defaults_sentinel_fills_all_ports() {
let mut raw = ConnectDecl::default();
raw.inputs.insert(String::new(), vec![".*RouteLevel.*".to_string()]);
let result = apply_connect_defaults(
raw, &ports(&["kbd", "pad"]), &ports(&["synth"]), None, None,
);
assert_eq!(result.inputs.get("kbd"), Some(&pats(&[".*RouteLevel.*"])));
assert_eq!(result.inputs.get("pad"), Some(&pats(&[".*RouteLevel.*"])));
assert!(!result.inputs.contains_key(""));
}
#[test]
fn defaults_sentinel_takes_priority_over_global() {
let raw = decl_with_sentinel(".*RouteLevel.*");
let result = apply_connect_defaults(
raw, &ports(&["kbd"]), &ports(&[]), Some(".*Global.*"), None,
);
assert_eq!(result.inputs.get("kbd"), Some(&pats(&[".*RouteLevel.*"])));
}
#[test]
fn defaults_per_port_not_overridden_by_global() {
let mut raw = ConnectDecl::default();
raw.inputs.insert("kbd".into(), vec![".*PerPort.*".to_string()]);
let result = apply_connect_defaults(
raw, &ports(&["kbd"]), &ports(&[]), Some(".*Global.*"), None,
);
assert_eq!(result.inputs.get("kbd"), Some(&pats(&[".*PerPort.*"])));
}
#[test]
fn defaults_per_port_not_overridden_by_sentinel() {
let mut raw = ConnectDecl::default();
raw.inputs.insert("kbd".into(), vec![".*PerPort.*".to_string()]);
raw.inputs.insert(String::new(), vec![".*Sentinel.*".to_string()]);
let result = apply_connect_defaults(
raw, &ports(&["kbd", "pad"]), &ports(&[]), None, None,
);
assert_eq!(result.inputs.get("kbd"), Some(&pats(&[".*PerPort.*"])));
assert_eq!(result.inputs.get("pad"), Some(&pats(&[".*Sentinel.*"])));
}
#[test]
fn defaults_per_port_on_one_input_suppresses_global_for_other_inputs() {
let mut raw = ConnectDecl::default();
raw.inputs.insert("kbd".into(), vec![".*PerPort.*".to_string()]);
let result = apply_connect_defaults(
raw, &ports(&["kbd", "pad"]), &ports(&[]), Some(".*Global.*"), None,
);
assert_eq!(result.inputs.get("kbd"), Some(&pats(&[".*PerPort.*"])));
assert!(!result.inputs.contains_key("pad"), "global must not fill 'pad' when route has any connect pattern");
}
#[test]
fn defaults_per_port_on_one_output_suppresses_global_for_other_outputs() {
let mut raw = ConnectDecl::default();
raw.outputs.insert("synth".into(), vec![".*PerPort.*".to_string()]);
let result = apply_connect_defaults(
raw, &ports(&[]), &ports(&["synth", "drums"]), None, Some(".*Global.*"),
);
assert_eq!(result.outputs.get("synth"), Some(&pats(&[".*PerPort.*"])));
assert!(!result.outputs.contains_key("drums"), "global must not fill 'drums' when route has any connect pattern");
}
#[test]
fn defaults_sentinel_removed_from_final_map() {
let raw = decl_with_sentinel(".*Pat.*");
let result = apply_connect_defaults(raw, &ports(&["kbd"]), &ports(&[]), None, None);
assert!(!result.inputs.contains_key(""));
}
fn state_path(dir: &std::path::Path) -> std::path::PathBuf {
dir.join("state.json")
}
#[test]
fn save_state_then_load_state_roundtrips() {
let lua = Lua::new();
let dir = std::env::temp_dir().join("midi_daemon_test_route_state_roundtrip");
let _ = std::fs::remove_dir_all(&dir);
register_state_functions(&lua, &state_path(&dir), "test").unwrap();
lua.load(r#"save_state({ count = 3, label = "hi" })"#).exec().unwrap();
let loaded: LuaTable = lua.load("return load_state()").eval().unwrap();
assert_eq!(loaded.get::<i64>("count").unwrap(), 3);
assert_eq!(loaded.get::<String>("label").unwrap(), "hi");
let _ = std::fs::remove_dir_all(&dir);
}
#[test]
fn load_state_before_any_save_returns_empty_table() {
let lua = Lua::new();
let dir = std::env::temp_dir().join("midi_daemon_test_route_state_empty");
let _ = std::fs::remove_dir_all(&dir);
register_state_functions(&lua, &state_path(&dir), "test").unwrap();
let loaded: LuaTable = lua.load("return load_state()").eval().unwrap();
assert_eq!(loaded.pairs::<LuaValue, LuaValue>().count(), 0);
let _ = std::fs::remove_dir_all(&dir);
}
#[test]
fn save_state_overwrites_previous_value() {
let lua = Lua::new();
let dir = std::env::temp_dir().join("midi_daemon_test_route_state_overwrite");
let _ = std::fs::remove_dir_all(&dir);
register_state_functions(&lua, &state_path(&dir), "test").unwrap();
lua.load("save_state({ n = 1 })").exec().unwrap();
lua.load("save_state({ n = 2 })").exec().unwrap();
let loaded: LuaTable = lua.load("return load_state()").eval().unwrap();
assert_eq!(loaded.get::<i64>("n").unwrap(), 2);
let _ = std::fs::remove_dir_all(&dir);
}
#[test]
fn call_optional_hook_runs_defined_hook() {
let lua = Lua::new();
lua.load("ran = false; function on_startup() ran = true end").exec().unwrap();
call_optional_hook(&lua, "on_startup", "test");
assert!(lua.globals().get::<bool>("ran").unwrap());
}
#[test]
fn call_optional_hook_is_noop_when_undefined() {
let lua = Lua::new();
call_optional_hook(&lua, "on_startup", "test");
}
#[test]
fn call_optional_hook_swallows_errors_from_the_hook() {
let lua = Lua::new();
lua.load("function on_shutdown() error('boom') end").exec().unwrap();
call_optional_hook(&lua, "on_shutdown", "test");
}
#[test]
fn call_optional_hook_ignores_a_non_function_global_of_the_same_name() {
let lua = Lua::new();
lua.load("on_startup = 42").exec().unwrap();
call_optional_hook(&lua, "on_startup", "test");
}
#[test]
fn save_state_and_load_state_are_independent_of_hooks() {
let lua = Lua::new();
let dir = std::env::temp_dir().join("midi_daemon_test_route_state_via_hook");
let _ = std::fs::remove_dir_all(&dir);
std::fs::create_dir_all(&dir).unwrap();
std::fs::write(state_path(&dir), r#"{"restored": 99}"#).unwrap();
register_state_functions(&lua, &state_path(&dir), "test").unwrap();
lua.load(r"
restored_value = nil
function on_startup()
local state = load_state()
restored_value = state.restored
end
").exec().unwrap();
call_optional_hook(&lua, "on_startup", "test");
assert_eq!(lua.globals().get::<i64>("restored_value").unwrap(), 99);
let _ = std::fs::remove_dir_all(&dir);
}
#[test]
fn obs_table_absent_returns_default() {
let lua = Lua::new();
let tbl = lua.create_table().unwrap();
let decl = obs_from_lua_table(&tbl);
assert_eq!(decl, ObsDecl::default());
}
#[test]
fn obs_connections_string_shorthand() {
let lua = Lua::new();
let tbl = lua.create_table().unwrap();
let obs_tbl = lua.create_table().unwrap();
obs_tbl.set("connections", "main").unwrap();
tbl.set("obs", obs_tbl).unwrap();
let decl = obs_from_lua_table(&tbl);
assert_eq!(decl.connections, vec!["main"]);
}
#[test]
fn obs_connections_array() {
let lua = Lua::new();
let tbl = lua.create_table().unwrap();
let obs_tbl = lua.create_table().unwrap();
obs_tbl.set("connections", lua_array(&lua, &["main", "backup"])).unwrap();
tbl.set("obs", obs_tbl).unwrap();
let decl = obs_from_lua_table(&tbl);
assert_eq!(decl.connections, vec!["main", "backup"]);
}
#[test]
fn obs_table_present_without_connections_returns_empty() {
let lua = Lua::new();
let tbl = lua.create_table().unwrap();
let obs_tbl = lua.create_table().unwrap();
tbl.set("obs", obs_tbl).unwrap();
let decl = obs_from_lua_table(&tbl);
assert!(decl.connections.is_empty());
}
}