use anyhow::Context;
use mlua::prelude::*;
use mlua::LuaSerdeExt;
use std::path::Path;
pub fn load_json_state(lua: &Lua, path: &Path) -> LuaResult<LuaTable> {
let empty = || lua.create_table();
if !path.exists() {
return empty();
}
let text = match std::fs::read_to_string(path) {
Ok(t) => t,
Err(e) => {
tracing::warn!("load_json_state: read {}: {}", path.display(), e);
return empty();
}
};
let json: serde_json::Value = match serde_json::from_str(&text) {
Ok(v) => v,
Err(e) => {
tracing::warn!("load_json_state: parse {}: {}", path.display(), e);
return empty();
}
};
match lua.to_value(&json)? {
LuaValue::Table(t) => Ok(t),
_ => empty(),
}
}
pub fn save_json_state(lua: &Lua, path: &Path, table: &LuaTable) -> anyhow::Result<()> {
let json: serde_json::Value = lua
.from_value(LuaValue::Table(table.clone()))
.map_err(|e| anyhow::anyhow!("serialize Lua state table: {e}"))?;
let text = serde_json::to_string_pretty(&json).context("serialize JSON state")?;
if let Some(parent) = path.parent() {
std::fs::create_dir_all(parent)
.with_context(|| format!("create state dir {}", parent.display()))?;
}
std::fs::write(path, text).with_context(|| format!("write state {}", path.display()))?;
Ok(())
}
pub fn toml_table_to_lua(lua: &Lua, table: &toml::Table) -> LuaResult<LuaTable> {
let t = lua.create_table()?;
for (k, v) in table {
t.set(k.as_str(), toml_value_to_lua(lua, v)?)?;
}
Ok(t)
}
fn toml_value_to_lua(lua: &Lua, value: &toml::Value) -> LuaResult<LuaValue> {
match value {
toml::Value::String(s) => Ok(LuaValue::String(lua.create_string(s)?)),
toml::Value::Integer(i) => Ok(LuaValue::Integer(*i)),
toml::Value::Float(f) => Ok(LuaValue::Number(*f)),
toml::Value::Boolean(b) => Ok(LuaValue::Boolean(*b)),
toml::Value::Array(arr) => {
let t = lua.create_table()?;
for (i, v) in arr.iter().enumerate() {
t.set(i + 1, toml_value_to_lua(lua, v)?)?;
}
Ok(LuaValue::Table(t))
}
toml::Value::Table(tbl) => Ok(LuaValue::Table(toml_table_to_lua(lua, tbl)?)),
toml::Value::Datetime(dt) => Ok(LuaValue::String(lua.create_string(dt.to_string())?)),
}
}
pub fn midi_bytes_to_lua(lua: &Lua, bytes: &[u8]) -> LuaResult<LuaTable> {
let msg = lua.create_table()?;
match bytes {
[status, note, vel] if (0x90..=0x9F).contains(status) && *vel > 0 => {
msg.set("type", "note_on")?;
msg.set("channel", (status & 0x0F) + 1)?;
msg.set("note", *note)?;
msg.set("velocity", *vel)?;
}
[status, note, vel]
if (0x80..=0x8F).contains(status)
|| ((0x90..=0x9F).contains(status) && *vel == 0) =>
{
msg.set("type", "note_off")?;
msg.set("channel", (status & 0x0F) + 1)?;
msg.set("note", *note)?;
msg.set("velocity", *vel)?;
}
[status, cc, val] if (0xB0..=0xBF).contains(status) => {
msg.set("type", "cc")?;
msg.set("channel", (status & 0x0F) + 1)?;
msg.set("controller", *cc)?;
msg.set("value", *val)?;
}
[status, prog] if (0xC0..=0xCF).contains(status) => {
msg.set("type", "program_change")?;
msg.set("channel", (status & 0x0F) + 1)?;
msg.set("program", *prog)?;
}
[status, lsb, msb] if (0xE0..=0xEF).contains(status) => {
let value = ((i16::from(*msb) << 7) | i16::from(*lsb)) - 8192;
msg.set("type", "pitch_bend")?;
msg.set("channel", (status & 0x0F) + 1)?;
msg.set("value", value)?;
}
[0xF8] => {
msg.set("type", "clock")?;
}
[0xFA] => {
msg.set("type", "start")?;
}
[0xFC] => {
msg.set("type", "stop")?;
}
[0xFB] => {
msg.set("type", "continue")?;
}
_ => {
msg.set("type", "raw")?;
let data = lua.create_table()?;
for (i, b) in bytes.iter().enumerate() {
data.set(i + 1, *b)?;
}
msg.set("data", data)?;
}
}
Ok(msg)
}
pub fn lua_to_midi_bytes(msg: &LuaTable) -> LuaResult<Vec<u8>> {
let msg_type: String = msg.get("type")?;
match msg_type.as_str() {
"note_on" => {
let ch: u8 = msg.get::<u8>("channel")?.saturating_sub(1) & 0x0F;
let note: u8 = msg.get("note")?;
let vel: u8 = msg.get("velocity")?;
Ok(vec![0x90 | ch, note, vel])
}
"note_off" => {
let ch: u8 = msg.get::<u8>("channel")?.saturating_sub(1) & 0x0F;
let note: u8 = msg.get("note")?;
let vel: u8 = msg.get::<Option<u8>>("velocity")?.unwrap_or(0);
Ok(vec![0x80 | ch, note, vel])
}
"cc" => {
let ch: u8 = msg.get::<u8>("channel")?.saturating_sub(1) & 0x0F;
let cc: u8 = msg.get("controller")?;
let val: u8 = msg.get("value")?;
Ok(vec![0xB0 | ch, cc, val])
}
"program_change" => {
let ch: u8 = msg.get::<u8>("channel")?.saturating_sub(1) & 0x0F;
let prog: u8 = msg.get("program")?;
Ok(vec![0xC0 | ch, prog])
}
"pitch_bend" => {
let ch: u8 = msg.get::<u8>("channel")?.saturating_sub(1) & 0x0F;
let value: i16 = msg.get("value")?;
let v = (value + 8192).clamp(0, 16383).cast_unsigned();
let lsb = (v & 0x7F) as u8;
let msb = ((v >> 7) & 0x7F) as u8;
Ok(vec![0xE0 | ch, lsb, msb])
}
"clock" => Ok(vec![0xF8]),
"start" => Ok(vec![0xFA]),
"stop" => Ok(vec![0xFC]),
"continue" => Ok(vec![0xFB]),
"raw" => {
let data: LuaTable = msg.get("data")?;
let len = data.len()?;
let mut bytes = Vec::with_capacity(usize::try_from(len).unwrap_or(0));
for i in 1..=len {
bytes.push(data.get::<u8>(i)?);
}
Ok(bytes)
}
other => Err(LuaError::RuntimeError(format!(
"Unknown MIDI message type: {other}"
))),
}
}
pub fn osc_message_to_lua(lua: &Lua, address: &str, args: &[rosc::OscType]) -> LuaResult<LuaTable> {
let msg = lua.create_table()?;
msg.set("address", address)?;
let args_tbl = lua.create_table()?;
for (i, arg) in args.iter().enumerate() {
args_tbl.set(i + 1, osc_type_to_lua_value(lua, arg)?)?;
}
msg.set("args", args_tbl)?;
Ok(msg)
}
pub(crate) fn osc_type_to_lua_value(lua: &Lua, t: &rosc::OscType) -> LuaResult<LuaValue> {
use rosc::OscType as O;
match t {
O::Int(n) => Ok(LuaValue::Integer(i64::from(*n))),
O::Long(n) => Ok(LuaValue::Integer(*n)),
O::Float(f) => Ok(LuaValue::Number(f64::from(*f))),
O::Double(d) => Ok(LuaValue::Number(*d)),
O::String(s) => Ok(LuaValue::String(lua.create_string(s)?)),
O::Bool(b) => Ok(LuaValue::Boolean(*b)),
O::Nil => Ok(LuaValue::Nil),
O::Inf => Ok(LuaValue::Number(f64::INFINITY)),
O::Blob(b) => Ok(LuaValue::String(lua.create_string(b)?)),
O::Char(c) => Ok(LuaValue::String(lua.create_string(c.encode_utf8(&mut [0u8; 4]))?)),
O::Time(t) => Ok(LuaValue::Number(
f64::from(t.seconds) + f64::from(t.fractional) / (f64::from(u32::MAX) + 1.0),
)),
O::Color(c) => {
let tbl = lua.create_table()?;
tbl.set("r", c.red)?;
tbl.set("g", c.green)?;
tbl.set("b", c.blue)?;
tbl.set("a", c.alpha)?;
Ok(LuaValue::Table(tbl))
}
O::Midi(m) => {
let tbl = lua.create_table()?;
tbl.set("port", m.port)?;
tbl.set("status", m.status)?;
tbl.set("data1", m.data1)?;
tbl.set("data2", m.data2)?;
Ok(LuaValue::Table(tbl))
}
O::Array(a) => {
let tbl = lua.create_table()?;
for (i, item) in a.content.iter().enumerate() {
tbl.set(i + 1, osc_type_to_lua_value(lua, item)?)?;
}
Ok(LuaValue::Table(tbl))
}
}
}
pub fn lua_val_to_osc_type(v: &LuaValue) -> LuaResult<rosc::OscType> {
match v {
LuaValue::Integer(n) => i32::try_from(*n).map(rosc::OscType::Int).map_err(|_| {
LuaError::RuntimeError(format!(
"send_osc: integer {} is out of range for OSC Int32 ({} to {})",
n,
i32::MIN,
i32::MAX,
))
}),
LuaValue::Number(f) => {
if f.abs() > f64::from(f32::MAX) {
return Err(LuaError::RuntimeError(format!(
"send_osc: float {f} is out of range for OSC Float32; use a smaller value"
)));
}
#[allow(clippy::cast_possible_truncation)]
Ok(rosc::OscType::Float(*f as f32))
}
LuaValue::String(s) => Ok(rosc::OscType::String(
s.to_str().map_err(LuaError::external)?.to_string(),
)),
LuaValue::Boolean(b) => Ok(rosc::OscType::Bool(*b)),
LuaValue::Nil => Ok(rosc::OscType::Nil),
other => Err(LuaError::RuntimeError(format!(
"send_osc: cannot convert {} to an OSC type (use integer, number, string, boolean, or nil)",
other.type_name()
))),
}
}
#[cfg(test)]
#[allow(clippy::float_cmp)] mod tests {
use super::*;
fn lua() -> Lua {
Lua::new()
}
fn str_field(t: &LuaTable, k: &str) -> String {
t.get::<String>(k).unwrap()
}
fn u8_field(t: &LuaTable, k: &str) -> u8 {
t.get::<u8>(k).unwrap()
}
fn i16_field(t: &LuaTable, k: &str) -> i16 {
t.get::<i16>(k).unwrap()
}
#[test]
fn parse_note_on() {
let lua = lua();
let msg = midi_bytes_to_lua(&lua, &[0x90, 60, 100]).unwrap();
assert_eq!(str_field(&msg, "type"), "note_on");
assert_eq!(u8_field(&msg, "channel"), 1);
assert_eq!(u8_field(&msg, "note"), 60);
assert_eq!(u8_field(&msg, "velocity"), 100);
}
#[test]
fn parse_note_on_channel_16() {
let lua = lua();
let msg = midi_bytes_to_lua(&lua, &[0x9F, 48, 64]).unwrap();
assert_eq!(str_field(&msg, "type"), "note_on");
assert_eq!(u8_field(&msg, "channel"), 16);
assert_eq!(u8_field(&msg, "note"), 48);
}
#[test]
fn parse_note_off_via_8x_status() {
let lua = lua();
let msg = midi_bytes_to_lua(&lua, &[0x83, 60, 0]).unwrap();
assert_eq!(str_field(&msg, "type"), "note_off");
assert_eq!(u8_field(&msg, "channel"), 4);
assert_eq!(u8_field(&msg, "note"), 60);
assert_eq!(u8_field(&msg, "velocity"), 0);
}
#[test]
fn parse_note_off_via_note_on_vel0() {
let lua = lua();
let msg = midi_bytes_to_lua(&lua, &[0x90, 60, 0]).unwrap();
assert_eq!(str_field(&msg, "type"), "note_off");
assert_eq!(u8_field(&msg, "channel"), 1);
assert_eq!(u8_field(&msg, "note"), 60);
}
#[test]
fn parse_cc() {
let lua = lua();
let msg = midi_bytes_to_lua(&lua, &[0xB2, 7, 100]).unwrap();
assert_eq!(str_field(&msg, "type"), "cc");
assert_eq!(u8_field(&msg, "channel"), 3);
assert_eq!(u8_field(&msg, "controller"), 7);
assert_eq!(u8_field(&msg, "value"), 100);
}
#[test]
fn parse_cc_channel_bounds() {
let lua = lua();
let msg = midi_bytes_to_lua(&lua, &[0xBF, 64, 127]).unwrap();
assert_eq!(str_field(&msg, "type"), "cc");
assert_eq!(u8_field(&msg, "channel"), 16);
}
#[test]
fn parse_program_change() {
let lua = lua();
let msg = midi_bytes_to_lua(&lua, &[0xC0, 42]).unwrap();
assert_eq!(str_field(&msg, "type"), "program_change");
assert_eq!(u8_field(&msg, "channel"), 1);
assert_eq!(u8_field(&msg, "program"), 42);
}
#[test]
fn parse_pitch_bend_center() {
let lua = lua();
let msg = midi_bytes_to_lua(&lua, &[0xE0, 0x00, 0x40]).unwrap();
assert_eq!(str_field(&msg, "type"), "pitch_bend");
assert_eq!(u8_field(&msg, "channel"), 1);
assert_eq!(i16_field(&msg, "value"), 0);
}
#[test]
fn parse_pitch_bend_min() {
let lua = lua();
let msg = midi_bytes_to_lua(&lua, &[0xE0, 0x00, 0x00]).unwrap();
assert_eq!(i16_field(&msg, "value"), -8192);
}
#[test]
fn parse_pitch_bend_max() {
let lua = lua();
let msg = midi_bytes_to_lua(&lua, &[0xE0, 0x7F, 0x7F]).unwrap();
assert_eq!(i16_field(&msg, "value"), 8191);
}
#[test]
fn parse_clock() {
let lua = lua();
let msg = midi_bytes_to_lua(&lua, &[0xF8]).unwrap();
assert_eq!(str_field(&msg, "type"), "clock");
}
#[test]
fn parse_start() {
let lua = lua();
let msg = midi_bytes_to_lua(&lua, &[0xFA]).unwrap();
assert_eq!(str_field(&msg, "type"), "start");
}
#[test]
fn parse_stop() {
let lua = lua();
let msg = midi_bytes_to_lua(&lua, &[0xFC]).unwrap();
assert_eq!(str_field(&msg, "type"), "stop");
}
#[test]
fn parse_continue() {
let lua = lua();
let msg = midi_bytes_to_lua(&lua, &[0xFB]).unwrap();
assert_eq!(str_field(&msg, "type"), "continue");
}
#[test]
fn parse_raw_sysex_fallback() {
let lua = lua();
let msg = midi_bytes_to_lua(&lua, &[0xF0, 0x41, 0xF7]).unwrap();
assert_eq!(str_field(&msg, "type"), "raw");
let data: LuaTable = msg.get("data").unwrap();
assert_eq!(data.get::<u8>(1).unwrap(), 0xF0);
assert_eq!(data.get::<u8>(2).unwrap(), 0x41);
assert_eq!(data.get::<u8>(3).unwrap(), 0xF7);
}
#[test]
fn parse_raw_empty() {
let lua = lua();
let msg = midi_bytes_to_lua(&lua, &[]).unwrap();
assert_eq!(str_field(&msg, "type"), "raw");
let data: LuaTable = msg.get("data").unwrap();
assert_eq!(data.len().unwrap(), 0);
}
#[test]
fn encode_note_on() {
let lua = lua();
let t = lua.create_table().unwrap();
t.set("type", "note_on").unwrap();
t.set("channel", 1u8).unwrap();
t.set("note", 60u8).unwrap();
t.set("velocity", 100u8).unwrap();
assert_eq!(lua_to_midi_bytes(&t).unwrap(), vec![0x90, 60, 100]);
}
#[test]
fn encode_note_on_channel_16() {
let lua = lua();
let t = lua.create_table().unwrap();
t.set("type", "note_on").unwrap();
t.set("channel", 16u8).unwrap();
t.set("note", 48u8).unwrap();
t.set("velocity", 64u8).unwrap();
assert_eq!(lua_to_midi_bytes(&t).unwrap(), vec![0x9F, 48, 64]);
}
#[test]
fn encode_note_off() {
let lua = lua();
let t = lua.create_table().unwrap();
t.set("type", "note_off").unwrap();
t.set("channel", 1u8).unwrap();
t.set("note", 60u8).unwrap();
t.set("velocity", 0u8).unwrap();
assert_eq!(lua_to_midi_bytes(&t).unwrap(), vec![0x80, 60, 0]);
}
#[test]
fn encode_note_off_optional_velocity() {
let lua = lua();
let t = lua.create_table().unwrap();
t.set("type", "note_off").unwrap();
t.set("channel", 1u8).unwrap();
t.set("note", 60u8).unwrap();
let bytes = lua_to_midi_bytes(&t).unwrap();
assert_eq!(bytes[0], 0x80);
assert_eq!(bytes[2], 0);
}
#[test]
fn encode_cc() {
let lua = lua();
let t = lua.create_table().unwrap();
t.set("type", "cc").unwrap();
t.set("channel", 2u8).unwrap();
t.set("controller", 7u8).unwrap();
t.set("value", 100u8).unwrap();
assert_eq!(lua_to_midi_bytes(&t).unwrap(), vec![0xB1, 7, 100]);
}
#[test]
fn encode_program_change() {
let lua = lua();
let t = lua.create_table().unwrap();
t.set("type", "program_change").unwrap();
t.set("channel", 1u8).unwrap();
t.set("program", 42u8).unwrap();
assert_eq!(lua_to_midi_bytes(&t).unwrap(), vec![0xC0, 42]);
}
#[test]
fn encode_pitch_bend_center() {
let lua = lua();
let t = lua.create_table().unwrap();
t.set("type", "pitch_bend").unwrap();
t.set("channel", 1u8).unwrap();
t.set("value", 0i16).unwrap();
assert_eq!(lua_to_midi_bytes(&t).unwrap(), vec![0xE0, 0x00, 0x40]);
}
#[test]
fn encode_pitch_bend_min() {
let lua = lua();
let t = lua.create_table().unwrap();
t.set("type", "pitch_bend").unwrap();
t.set("channel", 1u8).unwrap();
t.set("value", -8192i16).unwrap();
assert_eq!(lua_to_midi_bytes(&t).unwrap(), vec![0xE0, 0x00, 0x00]);
}
#[test]
fn encode_pitch_bend_max() {
let lua = lua();
let t = lua.create_table().unwrap();
t.set("type", "pitch_bend").unwrap();
t.set("channel", 1u8).unwrap();
t.set("value", 8191i16).unwrap();
assert_eq!(lua_to_midi_bytes(&t).unwrap(), vec![0xE0, 0x7F, 0x7F]);
}
#[test]
fn encode_clock() {
let lua = lua();
let t = lua.create_table().unwrap();
t.set("type", "clock").unwrap();
assert_eq!(lua_to_midi_bytes(&t).unwrap(), vec![0xF8]);
}
#[test]
fn encode_start() {
let lua = lua();
let t = lua.create_table().unwrap();
t.set("type", "start").unwrap();
assert_eq!(lua_to_midi_bytes(&t).unwrap(), vec![0xFA]);
}
#[test]
fn encode_stop() {
let lua = lua();
let t = lua.create_table().unwrap();
t.set("type", "stop").unwrap();
assert_eq!(lua_to_midi_bytes(&t).unwrap(), vec![0xFC]);
}
#[test]
fn encode_continue() {
let lua = lua();
let t = lua.create_table().unwrap();
t.set("type", "continue").unwrap();
assert_eq!(lua_to_midi_bytes(&t).unwrap(), vec![0xFB]);
}
#[test]
fn encode_raw() {
let lua = lua();
let data = lua.create_table().unwrap();
data.set(1, 0xF0u8).unwrap();
data.set(2, 0x41u8).unwrap();
data.set(3, 0xF7u8).unwrap();
let t = lua.create_table().unwrap();
t.set("type", "raw").unwrap();
t.set("data", data).unwrap();
assert_eq!(lua_to_midi_bytes(&t).unwrap(), vec![0xF0, 0x41, 0xF7]);
}
#[test]
fn encode_unknown_type_is_error() {
let lua = lua();
let t = lua.create_table().unwrap();
t.set("type", "sysex").unwrap();
assert!(lua_to_midi_bytes(&t).is_err());
}
fn roundtrip(input: &[u8]) -> Vec<u8> {
let lua = lua();
let table = midi_bytes_to_lua(&lua, input).unwrap();
lua_to_midi_bytes(&table).unwrap()
}
#[test]
fn roundtrip_note_on() {
assert_eq!(roundtrip(&[0x90, 60, 100]), vec![0x90, 60, 100]);
}
#[test]
fn roundtrip_note_on_all_channels() {
for ch in 0u8..16 {
let input = [0x90 | ch, 64, 80];
assert_eq!(roundtrip(&input), input.to_vec(), "channel {}", ch + 1);
}
}
#[test]
fn roundtrip_note_off() {
assert_eq!(roundtrip(&[0x80, 60, 0]), vec![0x80, 60, 0]);
}
#[test]
fn roundtrip_cc() {
assert_eq!(roundtrip(&[0xB3, 7, 100]), vec![0xB3, 7, 100]);
}
#[test]
fn roundtrip_program_change() {
assert_eq!(roundtrip(&[0xC5, 42]), vec![0xC5, 42]);
}
#[test]
fn roundtrip_pitch_bend_center() {
assert_eq!(roundtrip(&[0xE0, 0x00, 0x40]), vec![0xE0, 0x00, 0x40]);
}
#[test]
fn roundtrip_pitch_bend_min() {
assert_eq!(roundtrip(&[0xE0, 0x00, 0x00]), vec![0xE0, 0x00, 0x00]);
}
#[test]
fn roundtrip_pitch_bend_max() {
assert_eq!(roundtrip(&[0xE0, 0x7F, 0x7F]), vec![0xE0, 0x7F, 0x7F]);
}
#[test]
fn roundtrip_clock() {
assert_eq!(roundtrip(&[0xF8]), vec![0xF8]);
}
#[test]
fn roundtrip_start_stop_continue() {
assert_eq!(roundtrip(&[0xFA]), vec![0xFA]);
assert_eq!(roundtrip(&[0xFC]), vec![0xFC]);
assert_eq!(roundtrip(&[0xFB]), vec![0xFB]);
}
#[test]
fn load_json_state_missing_file_returns_empty_table() {
let lua = lua();
let path = std::path::PathBuf::from("/tmp/midi_daemon_test_state_nonexistent_xyz/state.json");
let t = load_json_state(&lua, &path).unwrap();
assert_eq!(t.len().unwrap(), 0);
}
#[test]
fn save_then_load_json_state_roundtrips() {
let lua = lua();
let dir = std::env::temp_dir().join("midi_daemon_test_state_roundtrip");
let _ = std::fs::remove_dir_all(&dir);
let path = dir.join("state.json");
let table = lua.create_table().unwrap();
table.set("bpm", 128.5).unwrap();
table.set("running", true).unwrap();
table.set("label", "hello").unwrap();
save_json_state(&lua, &path, &table).unwrap();
let loaded = load_json_state(&lua, &path).unwrap();
assert_eq!(loaded.get::<f64>("bpm").unwrap(), 128.5);
assert!(loaded.get::<bool>("running").unwrap());
assert_eq!(loaded.get::<String>("label").unwrap(), "hello");
let _ = std::fs::remove_dir_all(&dir);
}
#[test]
fn save_json_state_creates_parent_directory() {
let lua = lua();
let dir = std::env::temp_dir().join("midi_daemon_test_state_mkdir_parent");
let _ = std::fs::remove_dir_all(&dir);
let path = dir.join("state.json");
let table = lua.create_table().unwrap();
table.set("x", 1).unwrap();
save_json_state(&lua, &path, &table).unwrap();
assert!(path.exists());
let _ = std::fs::remove_dir_all(&dir);
}
#[test]
fn load_json_state_corrupt_file_returns_empty_table() {
let lua = lua();
let dir = std::env::temp_dir().join("midi_daemon_test_state_corrupt");
let _ = std::fs::remove_dir_all(&dir);
std::fs::create_dir_all(&dir).unwrap();
let path = dir.join("state.json");
std::fs::write(&path, "not valid json {{{").unwrap();
let t = load_json_state(&lua, &path).unwrap();
assert_eq!(t.len().unwrap(), 0);
let _ = std::fs::remove_dir_all(&dir);
}
}