diff --git a/Cargo.toml b/Cargo.toml index 7c91d37..6922a48 100644 --- a/Cargo.toml +++ b/Cargo.toml @@ -1,3 +1,3 @@ [workspace] -members = ["cli", "daemon", "libcompat/udev", "libs/ipc", "libs/kobject-uevent", "libs/rule"] +members = ["cli", "daemon", "libcompat/udev", "libs/db", "libs/ipc", "libs/kobject-uevent", "libs/rule"] resolver = "3" diff --git a/daemon/Cargo.toml b/daemon/Cargo.toml index 668dd28..8553ccd 100644 --- a/daemon/Cargo.toml +++ b/daemon/Cargo.toml @@ -15,3 +15,4 @@ serde_json = "1" tokio = { version = "1", features = ["rt", "rt-multi-thread", "macros", "fs", "signal", "sync"] } tracing = "0.1" tracing-subscriber = { version = "0.3", features = ["env-filter"] } +db = { path = "../libs/db" } diff --git a/daemon/src/device.rs b/daemon/src/device.rs index 06004e5..fd0a7ac 100644 --- a/daemon/src/device.rs +++ b/daemon/src/device.rs @@ -1,3 +1,4 @@ +use db::{Db, DbSection}; use kobject_uevent::KobjectUevent; use rule::{engine::Engine, parser::When}; use rustc_hash::FxHashMap; @@ -7,52 +8,57 @@ use std::{ path::{Path, PathBuf}, sync::{Arc, LazyLock, RwLock}, }; -use tokio::sync::Mutex; +use tokio::sync::{Mutex, broadcast}; static DEVICES: LazyLock>> = LazyLock::new(|| Default::default()); +static DB: LazyLock> = LazyLock::new(|| Db::open(db::PERSIST_PATH)); +static NOTIFY: LazyLock> = LazyLock::new(|| broadcast::channel(128).0); pub type ArcDevice = Arc; #[derive(Debug)] pub struct Device { lock: Mutex<()>, + properties: DbSection, engine: Engine, } impl Device { - pub fn new() -> Self { - let engine = Engine::new(); - engine.set_property("DEVROOT".into(), dev_root().to_string_lossy().into()); + pub fn new(name: &str) -> Self { + let engine = Engine::with_properties(DB.section(name)); + engine + .properties + .insert("DEVROOT", dev_root().to_string_lossy().into()); Self { lock: Mutex::new(()), + properties: DB.section(name), engine, } } fn apply_uevent(&self, ev: &KobjectUevent) { for (k, v) in ev.inner().iter() { - self.engine.set_property(k.into(), v.into()); + self.engine.properties.insert(k, v.into()); } self.update_derived_properties(); } fn update_derived_properties(&self) { - if let Some(devname) = self.engine.get_property("DEVNAME") { - self.engine.set_property( - "DEVNODE".into(), - format!("{}/{devname}", dev_root().display()), - ); + if let Some(devname) = self.engine.properties.get("DEVNAME") { + self.engine + .properties + .insert("DEVNODE", format!("{}/{devname}", dev_root().display())); } } async fn update_devnode_creds(&self) { - let Some(devnode) = self.engine.get_property("DEVNODE") else { + let Some(devnode) = self.engine.properties.get("DEVNODE") else { return; }; - let uid = crate::util::uid_by_string(self.engine.get_property("OWNER").as_deref()); - let gid = crate::util::gid_by_string(self.engine.get_property("GROUP").as_deref()); - let mode = crate::util::parse_mode(self.engine.get_property("MODE").as_deref()); + let uid = crate::util::uid_by_string(self.engine.properties.get("OWNER").as_deref()); + let gid = crate::util::gid_by_string(self.engine.properties.get("GROUP").as_deref()); + let mode = crate::util::parse_mode(self.engine.properties.get("MODE").as_deref()); if let Err(err) = crate::util::set_ownership(&devnode, uid, gid).await { tracing::warn!("{devnode}: failed to set ownership: {err}"); @@ -67,7 +73,7 @@ impl Device { } async fn update_symlinks(&self) { - let Some(devnode) = self.engine.get_property("DEVNODE") else { + let Some(devnode) = self.engine.properties.get("DEVNODE") else { return; }; for i in self.engine.get_list("SYMLINKS") { @@ -104,6 +110,9 @@ pub async fn handle_uevent(ev: &KobjectUevent) { } None => (), } + if let Some(devpath) = ev.devpath() { + _ = NOTIFY.send(devpath.into()); + } } async fn add(ev: &KobjectUevent) { @@ -112,7 +121,7 @@ async fn add(ev: &KobjectUevent) { let Some(devpath) = ev.devpath() else { return; }; - let device = Device::new(); + let device = Device::new(devpath); device.apply_uevent(&ev); device.engine.exec(&crate::rules::get()).await; @@ -140,6 +149,7 @@ async fn remove(ev: &KobjectUevent) { device.remove_symlinks().await; + device.properties.clear(); DEVICES.write().unwrap().remove(devpath); } @@ -221,6 +231,19 @@ pub async fn wait_for_idle() { } } +pub async fn wait_for_event() -> Option { + let mut rx = NOTIFY.subscribe(); + rx.recv().await.ok() +} + +pub async fn get_device(name: &str) -> DbSection { + DB.section(name) +} + +pub fn sync_db() { + DB.save(); +} + fn dev_root() -> PathBuf { std::env::var("LXDEVICED_DEV_ROOT") .as_deref() diff --git a/daemon/src/ipc.rs b/daemon/src/ipc.rs index bd462c5..5dd1d9d 100644 --- a/daemon/src/ipc.rs +++ b/daemon/src/ipc.rs @@ -3,6 +3,7 @@ use ipc::{ server::{Connection, Listener}, }; use rule::parser::When; +use rustc_hash::FxHashMap; use serde::Serialize; use std::pin::Pin; @@ -35,6 +36,8 @@ async fn handle_connection(mut connection: Connection) -> std::io::Result<()> { Request::ReloadRules => handler(reload_rules()), Request::ListDevices(when) => handler(list_devices(when)), Request::WaitForIdle => handler(wait_for_idle()), + Request::WaitForEvent => handler(wait_for_event()), + Request::GetDevice(name) => handler(get_device(name)), }; let reply = handler.await; connection.send(reply).await?; @@ -63,6 +66,16 @@ async fn wait_for_idle() -> Result<(), Error> { Ok(()) } +async fn wait_for_event() -> Result { + crate::device::wait_for_event() + .await + .ok_or(Error::PermissionDenied) +} + +async fn get_device(name: String) -> Result, Error> { + Ok(crate::device::get_device(&name).await.dump()) +} + fn handler( fut: impl Future> + Send + 'static, ) -> Pin> + Send>> { diff --git a/daemon/src/uevent.rs b/daemon/src/uevent.rs index bd33e05..2fd5ac6 100644 --- a/daemon/src/uevent.rs +++ b/daemon/src/uevent.rs @@ -15,6 +15,7 @@ pub fn start_monitor() -> std::io::Result<()> { crate::device::update_devnode(&ev).await; tokio::spawn(async move { crate::device::handle_uevent(&ev).await; + crate::device::sync_db(); }); } }); diff --git a/libcompat/udev/Cargo.toml b/libcompat/udev/Cargo.toml index 15867c2..c85f388 100644 --- a/libcompat/udev/Cargo.toml +++ b/libcompat/udev/Cargo.toml @@ -10,4 +10,6 @@ crate-type = ["cdylib"] ipc = { path = "../../libs/ipc" } libc = "0.2" rustc-hash = "2" +db = { path = "../../libs/db" } +fnmatch-regex = "0.3" rule = { path = "../../libs/rule" } diff --git a/libcompat/udev/src/lib.rs b/libcompat/udev/src/lib.rs index 17c708f..cdbdfbf 100644 --- a/libcompat/udev/src/lib.rs +++ b/libcompat/udev/src/lib.rs @@ -1,26 +1,28 @@ +mod util; + +use crate::util::FfiStringPool; +use db::Db; use ipc::client::Client; -use rule::parser::{Str, When}; +use rule::{ + engine::Engine, + parser::{Str, When}, +}; use rustc_hash::FxHashMap; use std::{ - ffi::{CStr, CString}, - str::FromStr, + ffi::CStr, + os::fd::AsRawFd, sync::atomic::{self, AtomicUsize}, }; #[derive(Debug)] pub struct Udev { - client: Client, userdata: *mut u8, refcount: AtomicUsize, } #[unsafe(no_mangle)] pub unsafe extern "C" fn udev_new() -> *mut Udev { - let Ok(client) = Client::connect() else { - return std::ptr::null_mut(); - }; Box::into_raw(Box::new(Udev { - client, userdata: std::ptr::null_mut(), refcount: AtomicUsize::new(1), })) @@ -61,8 +63,8 @@ pub unsafe extern "C" fn udev_set_userdata(udev: *mut Udev, userdata: *mut u8) { #[derive(Debug)] pub struct UdevListEntry { - key: CString, - value: CString, + key: *const libc::c_char, + value: *const libc::c_char, next: Option>, } @@ -87,7 +89,7 @@ pub unsafe extern "C" fn udev_list_entry_get_by_name( unsafe { let name = CStr::from_ptr(name); while !entry.is_null() { - if (*entry).key == name { + if CStr::from_ptr((*entry).key) == name { return entry; } entry = udev_list_entry_get_next(entry); @@ -100,14 +102,14 @@ pub unsafe extern "C" fn udev_list_entry_get_by_name( pub unsafe extern "C" fn udev_list_entry_get_name( entry: *const UdevListEntry, ) -> *const libc::c_char { - unsafe { (*entry).key.as_ptr() } + unsafe { (*entry).key } } #[unsafe(no_mangle)] pub unsafe extern "C" fn udev_list_entry_get_value( entry: *const UdevListEntry, ) -> *const libc::c_char { - unsafe { (*entry).value.as_ptr() } + unsafe { (*entry).value } } #[derive(Debug)] @@ -115,40 +117,33 @@ pub struct UdevDevice { refcount: AtomicUsize, udev: *mut Udev, properties: UdevListEntry, - syspath: Option, + syspath: Option<*const libc::c_char>, devlinks: UdevListEntry, tags: UdevListEntry, sysattr: UdevListEntry, - action: Option, seqnum: libc::c_ulonglong, - sysattr_cache: FxHashMap, + strings: FfiStringPool, } impl UdevDevice { fn new(udev: *mut Udev, properties: FxHashMap) -> Self { - let syspath_str = properties.get("DEVPATH").map(|x| format!("/sys/{x}")); - let syspath = syspath_str - .as_deref() - .map(|x| CString::from_str(x).unwrap_or_default()); - let devlinks = properties - .get("SYMLINKS") - .map(|x| x.as_str()) - .unwrap_or_default() - .split('\0') - .filter(|x| !x.is_empty()) - .map(|x| x.to_string()) - .collect(); - let devlinks = vec_to_list(&devlinks); - let tags = properties - .get("TAGS") - .map(|x| x.as_str()) - .unwrap_or_default() - .split('\0') - .filter(|x| !x.is_empty()) - .map(|x| x.to_string()) - .collect(); - let tags = vec_to_list(&tags); + let mut strings = FfiStringPool::new(); + let syspath = properties.get("DEVPATH").map(|x| format!("/sys/{x}")); + let devlinks = util::parse_zero_separated_list( + properties + .get("SYMLINKS") + .map(|x| x.as_str()) + .unwrap_or_default(), + ); + let devlinks = vec_to_list(&mut strings, &devlinks); + let tags = util::parse_zero_separated_list( + properties + .get("TAGS") + .map(|x| x.as_str()) + .unwrap_or_default(), + ); + let tags = vec_to_list(&mut strings, &tags); let mut sysattr = Vec::new(); - if let Some(syspath_str) = &syspath_str + if let Some(syspath_str) = &syspath && let Ok(read_dir) = std::fs::read_dir(syspath_str) { for i in read_dir { @@ -158,18 +153,18 @@ impl UdevDevice { sysattr.push(i.file_name().to_string_lossy().to_string()); } } - let sysattr = vec_to_list(&sysattr); + let sysattr = vec_to_list(&mut strings, &sysattr); + let properties = *hashmap_to_list(&mut strings, &properties); Self { refcount: AtomicUsize::new(1), - udev, - properties: *hashmap_to_list(&properties), - syspath, + udev: unsafe { udev_ref(udev) }, + properties, + syspath: syspath.map(|x| strings.insert_rust(x).as_ptr()), devlinks: *devlinks, tags: *tags, sysattr: *sysattr, - action: None, seqnum: 0, - sysattr_cache: Default::default(), + strings, } } } @@ -183,19 +178,58 @@ impl Drop for UdevDevice { #[unsafe(no_mangle)] pub unsafe extern "C" fn udev_device_new_from_syspath( - _udev: *mut Udev, - _syspath: *const libc::c_char, + udev: *mut Udev, + syspath: *const libc::c_char, ) -> *mut UdevDevice { - std::ptr::null_mut() + unsafe { + let Ok(syspath) = CStr::from_ptr(syspath).to_str() else { + return std::ptr::null_mut(); + }; + let Some(devpath) = syspath.strip_prefix("/sys") else { + return std::ptr::null_mut(); + }; + udev_device_new_from_lxdeviced_name(udev, devpath.into()) + } +} + +fn udev_device_new_from_lxdeviced_name(udev: *mut Udev, name: String) -> *mut UdevDevice { + let Ok(mut client) = Client::connect() else { + return std::ptr::null_mut(); + }; + let Ok(properties) = client.get_device(name) else { + return std::ptr::null_mut(); + }; + Box::into_raw(Box::new(UdevDevice::new(udev, properties))) } #[unsafe(no_mangle)] pub unsafe extern "C" fn udev_device_new_from_devnum( - _udev: *mut Udev, - _type: libc::c_char, - _devnum: libc::dev_t, + udev: *mut Udev, + ty: libc::c_char, + devnum: libc::dev_t, ) -> *mut UdevDevice { - std::ptr::null_mut() + let block = When::Match(Str::Direct("SUBSYSTEM".into()), Str::Direct("block".into())); + let ty = match ty { + b'b' => block, + b'c' => When::Not(Box::new(block)), + _ => return std::ptr::null_mut(), + }; + let major = libc::major(devnum); + let minor = libc::minor(devnum); + let major = When::Match(Str::Direct("MAJOR".into()), Str::Direct(major.to_string())); + let minor = When::Match(Str::Direct("MINOR".into()), Str::Direct(minor.to_string())); + let devnum = When::All(vec![major, minor]); + let when = When::All(vec![ty, devnum]); + let Ok(mut client) = Client::connect() else { + return std::ptr::null_mut(); + }; + let Ok(result) = client.list_devices(when) else { + return std::ptr::null_mut(); + }; + let Some(devname) = result.first() else { + return std::ptr::null_mut(); + }; + udev_device_new_from_lxdeviced_name(udev, devname.into()) } #[unsafe(no_mangle)] @@ -277,13 +311,7 @@ pub unsafe extern "C" fn udev_device_get_devtype(dev: *const UdevDevice) -> *con #[unsafe(no_mangle)] pub unsafe extern "C" fn udev_device_get_syspath(dev: *const UdevDevice) -> *const libc::c_char { - unsafe { - (*dev) - .syspath - .as_ref() - .map(|x| x.as_ptr()) - .unwrap_or_default() - } + unsafe { (*dev).syspath.as_ref().copied().unwrap_or_default() } } #[unsafe(no_mangle)] @@ -292,12 +320,15 @@ pub unsafe extern "C" fn udev_device_get_sysname(dev: *const UdevDevice) -> *con if dev.is_null() { return std::ptr::null(); } - let Some(syspath) = (*dev).syspath.as_ref() else { + let Some(syspath) = (*dev).syspath else { return std::ptr::null(); }; - let split_pos = syspath.as_bytes().iter().rposition(|x| *x == b'/'); + let split_pos = CStr::from_ptr(syspath) + .to_bytes() + .iter() + .rposition(|x| *x == b'/'); let offset = split_pos.map(|x| x + 1).unwrap_or_default(); - syspath.as_ptr().add(offset) + syspath.add(offset) } } @@ -372,7 +403,13 @@ pub unsafe extern "C" fn udev_device_get_property_value( dev: *const UdevDevice, key: *const libc::c_char, ) -> *const libc::c_char { - unsafe { udev_list_entry_get_value(udev_list_entry_get_by_name(&(*dev).properties, key)) } + unsafe { + let entry = udev_list_entry_get_by_name(&(*dev).properties, key); + if entry.is_null() { + return std::ptr::null_mut(); + } + udev_list_entry_get_value(entry) + } } #[unsafe(no_mangle)] @@ -398,13 +435,7 @@ pub unsafe extern "C" fn udev_device_get_devnum(dev: *const UdevDevice) -> libc: #[unsafe(no_mangle)] pub unsafe extern "C" fn udev_device_get_action(dev: *const UdevDevice) -> *const libc::c_char { - unsafe { - (*dev) - .action - .as_ref() - .map(|x| x.as_ptr()) - .unwrap_or_default() - } + unsafe { udev_device_get_property_value(dev, c"ACTION".as_ptr()) } } #[unsafe(no_mangle)] @@ -429,26 +460,21 @@ pub unsafe extern "C" fn udev_device_get_sysattr_value( let Ok(sysattr_str) = sysattr.to_str() else { return std::ptr::null(); }; - if let Some(val) = (*dev).sysattr_cache.get(sysattr) { - return val.as_ptr(); - } let Some(syspath) = (*dev).syspath.as_ref() else { return std::ptr::null(); }; - let Ok(mut path) = String::from_utf8(syspath.as_bytes().to_vec()) else { + let Ok(mut path) = String::from_utf8(CStr::from_ptr(*syspath).to_bytes().to_vec()) else { return std::ptr::null(); }; path.push('/'); path.push_str(sysattr_str); - let Ok(data) = std::fs::read_to_string(path) else { + let Ok(mut data) = std::fs::read_to_string(path) else { return std::ptr::null(); }; - let Ok(data) = CString::from_str(&data) else { - return std::ptr::null(); - }; - let ptr = data.as_ptr(); - (*dev).sysattr_cache.insert(CString::from(sysattr), data); - ptr + if data.ends_with("\n") { + data.pop(); + } + (*dev).strings.insert_rust(data).as_ptr() } } @@ -472,21 +498,12 @@ pub unsafe extern "C" fn udev_device_set_sysattr_value( Some(p) => p, None => return -libc::ENOENT, }; - let path = match syspath.to_str() { + let path = match CStr::from_ptr(*syspath).to_str() { Ok(p) => format!("{}/{}", p, sysattr), Err(_) => return -libc::ENOENT, }; match std::fs::write(&path, value) { - Ok(_) => { - if let Ok(cstr) = CString::new(value) { - let key = match CString::new(sysattr) { - Ok(k) => k, - Err(_) => return 0, - }; - dev.sysattr_cache.insert(key, cstr); - } - 0 - } + Ok(_) => 0, Err(e) => match e.raw_os_error() { Some(code) => -code, None => -libc::EIO, @@ -519,6 +536,8 @@ pub unsafe extern "C" fn udev_device_has_current_tag( pub struct UdevMonitor { udev: *mut Udev, refcount: AtomicUsize, + client: Client, + conditions: Vec, } impl Drop for UdevMonitor { fn drop(&mut self) { @@ -560,20 +579,30 @@ pub unsafe extern "C" fn udev_monitor_new_from_netlink( name: *const libc::c_char, ) -> *mut UdevMonitor { unsafe { - let name = CStr::from_ptr(name); - if name != c"udev" { + if CStr::from_ptr(name) != c"udev" { return std::ptr::null_mut(); } + let Ok(client) = Client::connect() else { + return std::ptr::null_mut(); + }; Box::into_raw(Box::new(UdevMonitor { - udev, + udev: udev_ref(udev), refcount: AtomicUsize::new(1), + client, + conditions: Vec::new(), })) } } #[unsafe(no_mangle)] -pub unsafe extern "C" fn udev_monitor_enable_receiving(_: *const UdevMonitor) -> libc::c_int { - -libc::EPERM +pub unsafe extern "C" fn udev_monitor_enable_receiving(monitor: *mut UdevMonitor) -> libc::c_int { + unsafe { + let monitor = &mut *monitor; + match monitor.client.send(&ipc::repr::Request::WaitForEvent) { + Ok(()) => 0, + Err(err) => -err.raw_os_error().unwrap_or(libc::EPERM), + } + } } #[unsafe(no_mangle)] @@ -585,40 +614,94 @@ pub unsafe extern "C" fn udev_monitor_set_receive_buffer_size( } #[unsafe(no_mangle)] -pub unsafe extern "C" fn udev_monitor_get_fd(_: *const UdevMonitor) -> libc::c_int { - -1 +pub unsafe extern "C" fn udev_monitor_get_fd(monitor: *const UdevMonitor) -> libc::c_int { + unsafe { (*monitor).client.inner().as_raw_fd() } } #[unsafe(no_mangle)] -pub unsafe extern "C" fn udev_monitor_receive_device(_: *const UdevMonitor) -> *mut UdevDevice { - std::ptr::null_mut() +pub unsafe extern "C" fn udev_monitor_receive_device(monitor: *mut UdevMonitor) -> *mut UdevDevice { + unsafe { + let monitor = &mut *monitor; + let db = Db::open_in_memory(); + loop { + let result = monitor.client.recv::(); + udev_monitor_enable_receiving(monitor); + let Ok(Ok(result)) = result else { + return std::ptr::null_mut(); + }; + let Ok(device) = monitor.client.get_device(result) else { + return std::ptr::null_mut(); + }; + + let section = db.section("default"); + section.restore(device.clone()); + let engine = Engine::with_properties(section); + if !engine + .exec_condition(&When::All(monitor.conditions.clone())) + .unwrap_or_default() + { + return std::ptr::null_mut(); + } + + return Box::into_raw(Box::new(UdevDevice::new(monitor.udev, device))); + } + } } #[unsafe(no_mangle)] pub unsafe extern "C" fn udev_monitor_filter_add_match_subsystem_devtype( - _mon: *mut UdevMonitor, - _subsystem: *const libc::c_char, - _devtype: *const libc::c_char, + mon: *mut UdevMonitor, + subsystem: *const libc::c_char, + devtype: *const libc::c_char, ) -> libc::c_int { - -libc::EPERM + unsafe { + udev_monitor_match(mon, "SUBSYSTEM", subsystem, false); + udev_monitor_match(mon, "DEVTYPE", devtype, false); + 0 + } } #[unsafe(no_mangle)] pub unsafe extern "C" fn udev_monitor_filter_add_match_tag( - _mon: *mut UdevMonitor, - _tag: *const libc::c_char, + mon: *mut UdevMonitor, + tag: *const libc::c_char, ) -> libc::c_int { - -libc::EPERM + unsafe { udev_monitor_match(mon, "TAGS", tag, false) } } #[unsafe(no_mangle)] pub unsafe extern "C" fn udev_monitor_filter_update(_mon: *mut UdevMonitor) -> libc::c_int { - -libc::EPERM + 0 } #[unsafe(no_mangle)] -pub unsafe extern "C" fn udev_monitor_filter_remove(_mon: *mut UdevMonitor) -> libc::c_int { - -libc::EPERM +pub unsafe extern "C" fn udev_monitor_filter_remove(mon: *mut UdevMonitor) -> libc::c_int { + unsafe { + (*mon).conditions.clear(); + 0 + } +} + +unsafe fn udev_monitor_match( + this: *mut UdevMonitor, + key: &str, + data: *const libc::c_char, + not: bool, +) -> libc::c_int { + unsafe { + if data.is_null() { + return 0; + } + let Ok(data) = CStr::from_ptr(data).to_str() else { + return -libc::EINVAL; + }; + let mut when = When::Match(Str::Direct(key.into()), Str::Direct(data.into())); + if not { + when = When::Not(Box::new(when)); + } + (*this).conditions.push(when); + 0 + } } #[derive(Debug)] @@ -627,6 +710,7 @@ pub struct UdevEnumerate { refcount: AtomicUsize, filters: Vec, result: Option, + strings: FfiStringPool, } impl Drop for UdevEnumerate { fn drop(&mut self) { @@ -664,12 +748,15 @@ pub unsafe extern "C" fn udev_enumerate_get_udev(p: *mut UdevEnumerate) -> *mut #[unsafe(no_mangle)] pub unsafe extern "C" fn udev_enumerate_new(udev: *mut Udev) -> *mut UdevEnumerate { - Box::into_raw(Box::new(UdevEnumerate { - udev, - refcount: AtomicUsize::new(1), - filters: Vec::with_capacity(16), - result: None, - })) + unsafe { + Box::into_raw(Box::new(UdevEnumerate { + udev: udev_ref(udev), + refcount: AtomicUsize::new(1), + filters: Vec::with_capacity(16), + result: None, + strings: FfiStringPool::new(), + })) + } } #[unsafe(no_mangle)] @@ -735,7 +822,7 @@ pub unsafe extern "C" fn udev_enumerate_add_match_sysname( udev_enumerate: *mut UdevEnumerate, sysname: *const libc::c_char, ) -> libc::c_int { - unsafe { udev_enumerate_match(udev_enumerate, "SYSNAME", sysname, false) } + unsafe { udev_enumerate_match(udev_enumerate, "DEVNAME", sysname, false) } } #[unsafe(no_mangle)] @@ -774,13 +861,18 @@ pub unsafe extern "C" fn udev_enumerate_scan_devices( udev_enumerate: *mut UdevEnumerate, ) -> libc::c_int { unsafe { - let udev = (*udev_enumerate).udev; let when = When::All((*udev_enumerate).filters.clone()); - let result = match (*udev).client.list_devices(when) { + let Ok(mut client) = Client::connect() else { + return -libc::EPERM; + }; + let mut result = match client.list_devices(when) { Ok(x) => x, Err(_) => return -libc::EPERM, }; - (*udev_enumerate).result = Some(*vec_to_list(&result)); + for i in &mut result { + i.insert_str(0, "/sys"); + } + (*udev_enumerate).result = Some(*vec_to_list(&mut (*udev_enumerate).strings, &result)); 0 } } @@ -815,7 +907,8 @@ unsafe fn udev_enumerate_match( let Ok(data) = CStr::from_ptr(data).to_str() else { return -libc::EINVAL; }; - let mut when = When::Match(Str::Direct(key.into()), Str::Direct(data.into())); + let regex = fnmatch_regex::glob_to_regex_pattern(data).unwrap_or_default(); + let mut when = When::PartMatch(Str::Direct(key.into()), Str::Direct((®ex[1..]).into())); if not { when = When::Not(Box::new(when)); } @@ -906,16 +999,17 @@ pub unsafe extern "C" fn udev_util_encode_string( len as _ } -fn hashmap_to_list(hashmap: &FxHashMap) -> Box { +fn hashmap_to_list( + strings: &mut FfiStringPool, + hashmap: &FxHashMap, +) -> Box { let mut head = None; let mut tail = &mut head; for (key, value) in hashmap { - let key_cstr = CString::new(key.as_str()).unwrap(); - let value_cstr = CString::new(value.as_str()).unwrap(); let entry = Box::new(UdevListEntry { - key: key_cstr, - value: value_cstr, + key: strings.insert_rust(key.clone()).as_ptr(), + value: strings.insert_rust(value.clone()).as_ptr(), next: None, }); *tail = Some(entry); @@ -923,21 +1017,20 @@ fn hashmap_to_list(hashmap: &FxHashMap) -> Box { } head.unwrap_or(Box::new(UdevListEntry { - key: CString::new("").unwrap(), - value: CString::new("").unwrap(), + key: std::ptr::null(), + value: std::ptr::null(), next: None, })) } -fn vec_to_list(vec: &Vec) -> Box { +fn vec_to_list(strings: &mut FfiStringPool, vec: &Vec) -> Box { let mut head = None; let mut tail = &mut head; for s in vec { - let cstr = CString::new(s.as_str()).unwrap(); let entry = Box::new(UdevListEntry { - key: cstr, - value: CString::new("").unwrap(), + key: strings.insert_rust(s.clone()).as_ptr(), + value: strings.insert_rust("".into()).as_ptr(), next: None, }); *tail = Some(entry); @@ -946,8 +1039,8 @@ fn vec_to_list(vec: &Vec) -> Box { head.unwrap_or_else(|| { Box::new(UdevListEntry { - key: CString::new("").unwrap(), - value: CString::new("").unwrap(), + key: std::ptr::null(), + value: std::ptr::null(), next: None, }) }) diff --git a/libcompat/udev/src/util.rs b/libcompat/udev/src/util.rs new file mode 100644 index 0000000..12bc8a1 --- /dev/null +++ b/libcompat/udev/src/util.rs @@ -0,0 +1,28 @@ +use rustc_hash::FxHashMap; +use std::ffi::{CStr, CString}; + +#[derive(Debug)] +pub struct FfiStringPool { + strings: FxHashMap, +} +impl FfiStringPool { + pub fn new() -> Self { + Self { + strings: FxHashMap::default(), + } + } + + pub fn insert_rust(&mut self, s: String) -> &CStr { + let s = s.replace("\0", ""); + let entry = self.strings.entry(s); + let value = entry.or_insert_with_key(|key| CString::new(key.clone()).unwrap()); + unsafe { &*(value.as_c_str() as *const CStr) } + } +} + +pub fn parse_zero_separated_list(s: &str) -> Vec { + s.split('\0') + .filter(|x| !x.is_empty()) + .map(|x| x.to_string()) + .collect() +} diff --git a/libs/db/Cargo.toml b/libs/db/Cargo.toml new file mode 100644 index 0000000..b47914c --- /dev/null +++ b/libs/db/Cargo.toml @@ -0,0 +1,8 @@ +[package] +name = "db" +version = "0.1.0" +edition = "2024" + +[dependencies] +rustc-hash = "2" +serde_json = "1" diff --git a/libs/db/src/lib.rs b/libs/db/src/lib.rs new file mode 100644 index 0000000..66a8f4e --- /dev/null +++ b/libs/db/src/lib.rs @@ -0,0 +1,115 @@ +use rustc_hash::FxHashMap; +use std::{ + path::PathBuf, + sync::{Arc, RwLock}, +}; + +pub const PERSIST_PATH: &str = "/run/lxdeviced.json"; + +#[derive(Debug)] +pub struct Db { + path: Option, + map: RwLock>, +} +impl Db { + pub fn open_in_memory() -> Arc { + Arc::new(Self { + path: None, + map: Default::default(), + }) + } + + pub fn open(path: impl Into) -> Arc { + let path = path.into(); + let text = std::fs::read_to_string(&path).unwrap_or_default(); + let map = serde_json::from_str(&text).unwrap_or_default(); + Arc::new(Self { + path: Some(path), + map: RwLock::new(map), + }) + } + + pub fn section(self: &Arc, name: impl Into) -> DbSection { + DbSection { + name: name.into(), + db: self.clone(), + } + } + + pub fn save(&self) { + if let Some(path) = &self.path { + let text = serde_json::to_string(&*self.map.read().unwrap()).unwrap(); + _ = std::fs::write(&path, text); + } + } +} + +#[derive(Debug)] +pub struct DbSection { + name: String, + db: Arc, +} +impl DbSection { + pub fn db(&self) -> &Arc { + &self.db + } + + pub fn clear(&self) { + let mut map = self.db.map.write().unwrap(); + let mut keys = Vec::new(); + let prefix = self.real_key_name(""); + for key in map.keys() { + if key.starts_with(&prefix) { + keys.push(key.clone()); + } + } + for key in keys { + map.remove(&key); + } + } + + pub fn insert(&self, key: &str, val: String) -> Option { + let mut map = self.db.map.write().unwrap(); + map.insert(self.real_key_name(key), val) + } + + pub fn get(&self, key: &str) -> Option { + self.db + .map + .read() + .unwrap() + .get(&self.real_key_name(key)) + .cloned() + } + + pub fn remove(&self, key: &str) -> Option { + self.db + .map + .write() + .unwrap() + .remove(&self.real_key_name(key)) + } + + pub fn restore(&self, dump: FxHashMap) { + self.clear(); + for (k, v) in dump { + self.insert(&k, v); + } + } + + pub fn dump(&self) -> FxHashMap { + let mut result = FxHashMap::default(); + let prefix = self.real_key_name(""); + let total = self.db.map.read().unwrap(); + for (key, val) in total.iter() { + if let Some(realkey) = key.strip_prefix(&prefix) { + result.insert(realkey.into(), val.into()); + } + } + result + } + + fn real_key_name(&self, name: &str) -> String { + format!("{}\0{}", self.name, name) + } +} diff --git a/libs/ipc/Cargo.toml b/libs/ipc/Cargo.toml index f9d19a8..9ff11a5 100644 --- a/libs/ipc/Cargo.toml +++ b/libs/ipc/Cargo.toml @@ -9,3 +9,4 @@ serde = { version = "1", features = ["derive"] } serde_json = "1" thiserror = "2" tokio = { version = "1", features = ["net", "io-util"] } +rustc-hash = "2" diff --git a/libs/ipc/src/client.rs b/libs/ipc/src/client.rs index 0422abf..9062a27 100644 --- a/libs/ipc/src/client.rs +++ b/libs/ipc/src/client.rs @@ -3,6 +3,7 @@ use crate::{ repr::{Error, Request, de_reply}, }; use rule::parser::When; +use rustc_hash::FxHashMap; use serde::de::DeserializeOwned; use std::{ io::{Read, Write}, @@ -26,6 +27,10 @@ impl Client { }) } + pub fn inner(&self) -> &UnixStream { + &self.stream + } + pub fn send(&mut self, req: &Request) -> std::io::Result<()> { self.buf.clear(); serde_json::to_writer(&mut self.buf, req).unwrap(); @@ -61,6 +66,10 @@ impl Client { self.invoke(&Request::ListDevices(when)) } + pub fn get_device(&mut self, name: String) -> Result, Error> { + self.invoke(&Request::GetDevice(name)) + } + pub fn wait_for_idle(&mut self) -> Result<(), Error> { self.invoke(&Request::WaitForIdle) } diff --git a/libs/ipc/src/repr.rs b/libs/ipc/src/repr.rs index 1cccf86..96c2cd0 100644 --- a/libs/ipc/src/repr.rs +++ b/libs/ipc/src/repr.rs @@ -15,6 +15,12 @@ pub enum Request { #[serde(rename = "/org.semilabs.os/LxDeviced/WaitForIdle@LXDEVICED_0.1.0")] WaitForIdle, + + #[serde(rename = "/org.semilabs.os/LxDeviced/WaitForEvent@LXDEVICED_0.1.0")] + WaitForEvent, + + #[serde(rename = "/org.semilabs.os/LxDeviced/GetDevice@LXDEVICED_0.1.0")] + GetDevice(String), } pub fn ser_reply(result: Result) -> String { diff --git a/libs/rule/Cargo.toml b/libs/rule/Cargo.toml index 3ca360d..0cec742 100644 --- a/libs/rule/Cargo.toml +++ b/libs/rule/Cargo.toml @@ -12,3 +12,4 @@ serde = { version = "1", features = ["derive"] } thiserror = "2" tokio = { version = "1", features = ["process", "time"] } tracing = "0.1" +db = { path = "../db" } diff --git a/libs/rule/src/engine.rs b/libs/rule/src/engine.rs index 86dd21c..86838b8 100644 --- a/libs/rule/src/engine.rs +++ b/libs/rule/src/engine.rs @@ -2,12 +2,12 @@ use crate::{ Error, parser::{Action, OnError, Rule, RuleFile, Str, When}, }; +use db::DbSection; use regex::Regex; -use rustc_hash::FxHashMap; use std::{ path::Path, process::{ExitStatus, Stdio}, - sync::{Arc, RwLock}, + sync::Arc, time::Duration, }; use tokio::{io::AsyncReadExt, process::Child}; @@ -29,29 +29,25 @@ impl RuleSet { ); continue; }; - for file in read_dir { - let Ok(file) = file else { - tracing::warn!( - "Error while reading directory \"{}\" for rule files.", - dir.as_ref().display() - ); - break; - }; - if file.file_name().as_encoded_bytes().starts_with(b".") { + let mut files: Vec<_> = read_dir.filter_map(|x| x.ok()).map(|x| x.path()).collect(); + files.sort_unstable(); + for file in files { + if file + .file_name() + .unwrap() + .as_encoded_bytes() + .starts_with(b".") + { continue; } - let rule_file = match RuleFile::parse(file.path()) { + let rule_file = match RuleFile::parse(&file) { Ok(x) => x, Err(e) => { - tracing::warn!( - "Failed to parse rule file \"{}\": {}", - file.path().display(), - e, - ); + tracing::warn!("Failed to parse rule file \"{}\": {}", file.display(), e,); continue; } }; - tracing::debug!("Activating rule file \"{}\"", file.path().display()); + tracing::debug!("Activating rule file \"{}\"", file.display()); rule_files.push(rule_file); } } @@ -61,31 +57,20 @@ impl RuleSet { #[derive(Debug)] pub struct Engine { - properties: RwLock>, + pub properties: DbSection, command_policy: CommandPolicy, } impl Engine { - pub fn new() -> Self { + pub fn with_properties(properties: DbSection) -> Self { Self { - properties: Default::default(), + properties, command_policy: CommandPolicy::default(), } } - pub fn get_property(&self, key: &str) -> Option { - self.properties.read().unwrap().get(key).cloned() - } - - pub fn set_property(&self, key: String, val: String) { - self.properties.write().unwrap().insert(key, val); - } - - pub fn remove_property(&self, key: &str) { - self.properties.write().unwrap().remove(key); - } - pub fn get_list(&self, key: &str) -> Vec { - self.get_property(key) + self.properties + .get(key) .unwrap_or_default() .split('\0') .filter(|x| !x.is_empty()) @@ -93,13 +78,15 @@ impl Engine { .collect() } - pub fn list_add(&self, key: String, val: &str) { - let orig = self.get_property(&key).unwrap_or_default(); - self.set_property(key, format!("{orig}{val}\0")); + pub fn list_add(&self, key: &str, val: &str) { + let orig = self.properties.get(&key).unwrap_or_default(); + self.properties.insert(&key, format!("{orig}{val}\0")); } pub async fn exec(&self, rule_set: &RuleSet) { for rule_file in rule_set.0.iter() { + self.properties + .insert("RULEFILE", rule_file.path.display().to_string()); for rule in rule_file.rules.iter() { if let Err(err) = self.exec_rule(rule).await { if !matches!(rule.on_error, OnError::Continue) { @@ -149,7 +136,7 @@ impl Engine { let regex = self.bake_str(regex)?; let regex = format!("^{regex}$"); let regex = Regex::new(®ex).map_err(Error::Regex)?; - let Some(val) = self.get_property(&key) else { + let Some(val) = self.properties.get(&key) else { return Ok(false); }; Ok(regex.is_match(&val)) @@ -158,7 +145,7 @@ impl Engine { let key = self.bake_str(key)?; let regex = self.bake_str(regex)?; let regex = Regex::new(®ex).map_err(Error::Regex)?; - let Some(val) = self.get_property(&key) else { + let Some(val) = self.properties.get(&key) else { return Ok(false); }; Ok(regex.is_match(&val)) @@ -171,15 +158,15 @@ impl Engine { Action::Set(key, val) => { let key = self.bake_str(key)?; let val = self.bake_str(val)?; - self.set_property(key, val); + self.properties.insert(&key, val); } Action::Add(key, val) => { let key = self.bake_str(key)?; let val = self.bake_str(val)?; - self.list_add(key, &val); + self.list_add(&key, &val); } Action::Unset(key) => { - self.remove_property(&self.bake_str(key)?); + self.properties.remove(&self.bake_str(key)?); } Action::Command(cmd) => { let cmd = self.bake_str(cmd)?; @@ -189,7 +176,7 @@ impl Engine { let key = self.bake_str(key)?; let cmd = self.bake_str(cmd)?; let val = self.command_policy.run_provider(&cmd).await?; - self.set_property(key, val); + self.properties.insert(&key, val); } } Ok(()) @@ -216,7 +203,8 @@ impl Engine { current_property = None; } else if ch == '}' { let value = self - .get_property(property) + .properties + .get(property) .ok_or_else(|| Error::PropertyNotFound(property.clone()))?; result.push_str(&value); current_property = None; @@ -239,11 +227,6 @@ impl Engine { } } } -impl Default for Engine { - fn default() -> Self { - Self::new() - } -} #[derive(Debug, Clone)] pub struct CommandPolicy { diff --git a/misc/rules/0030-input.ron b/misc/rules/0030-input.ron new file mode 100644 index 0000000..e1a98ec --- /dev/null +++ b/misc/rules/0030-input.ron @@ -0,0 +1,65 @@ +[ + ( + when: All([ + Match("SUBSYSTEM", "input"), + ]), + do: [ + UseProvider("SCRIPT_INPUT_ID", f("echo $(dirname {RULEFILE})/../scripts/input_id")), + ], + on_error: Continue, + ), + ( + when: All([ + Match("SCRIPT_INPUT_ID", ".*"), + ]), + do: [ + UseProvider("ID_INPUT", f("{SCRIPT_INPUT_ID} MODE=ID_INPUT DEVPATH={DEVPATH}")), + ], + on_error: Continue, + ), + ( + when: All([ + Match("SCRIPT_INPUT_ID", ".*"), + ]), + do: [ + UseProvider("ID_INPUT_KEYBOARD", f("{SCRIPT_INPUT_ID} MODE=ID_INPUT_KEYBOARD DEVPATH={DEVPATH}")), + ], + on_error: Continue, + ), + ( + when: All([ + Match("SCRIPT_INPUT_ID", ".*"), + ]), + do: [ + UseProvider("ID_INPUT_MOUSE", f("{SCRIPT_INPUT_ID} MODE=ID_INPUT_MOUSE DEVPATH={DEVPATH}")), + ], + on_error: Continue, + ), + ( + when: All([ + Match("SCRIPT_INPUT_ID", ".*"), + ]), + do: [ + UseProvider("ID_INPUT_TOUCHPAD", f("{SCRIPT_INPUT_ID} MODE=ID_INPUT_TOUCHPAD DEVPATH={DEVPATH}")), + ], + on_error: Continue, + ), + ( + when: All([ + Match("SCRIPT_INPUT_ID", ".*"), + ]), + do: [ + UseProvider("ID_INPUT_TOUCHSCREEN", f("{SCRIPT_INPUT_ID} MODE=ID_INPUT_TOUCHSCREEN DEVPATH={DEVPATH}")), + ], + on_error: Continue, + ), + ( + when: All([ + Match("SCRIPT_INPUT_ID", ".*"), + ]), + do: [ + UseProvider("ID_INPUT_JOYSTICK", f("{SCRIPT_INPUT_ID} MODE=ID_INPUT_JOYSTICK DEVPATH={DEVPATH}")), + ], + on_error: Continue, + ), +] diff --git a/misc/scripts/input_id b/misc/scripts/input_id new file mode 100755 index 0000000..6984b7e --- /dev/null +++ b/misc/scripts/input_id @@ -0,0 +1,125 @@ +#!/bin/mksh + +# Kernel constant definitions. +EV_KEY=1 +EV_REL=2 +EV_ABS=3 +EV_FF=21 +BTN_MOUSE=272 +ASCII_KEYS="30 48 46 32 18 33 34 35 23 36 37 38 50 49 24 25 16 19 31 20 22 47 17 45 21 44" +INPUT_PROP_POINTER=0 +INPUT_PROP_DIRECT=1 + +# Helper functions +bitmap_contains() { + hexstr="$1" + bitnum="$2" + + char_from_right=$(( bitnum / 4 )) + bit_in_char=$(( bitnum % 4 )) + + chars_per_group=16 + group_index=$(( char_from_right / chars_per_group )) + char_in_group=$(( char_from_right % chars_per_group )) + + set -A groups $hexstr + n=${#groups[@]} + + idx=$(( n - 1 - group_index )) + if [ "$idx" -lt 0 ] || [ "$idx" -ge "$n" ]; then + return 1 + fi + + group=${groups[idx]} + + while [ ${#group} -lt $chars_per_group ]; do + group="0$group" + done + + pos=$(( chars_per_group - 1 - char_in_group )) + ch=${group:$pos:1} + + case "$ch" in + 0) v=0 ;; 1) v=1 ;; 2) v=2 ;; 3) v=3 ;; + 4) v=4 ;; 5) v=5 ;; 6) v=6 ;; 7) v=7 ;; + 8) v=8 ;; 9) v=9 ;; + a|A) v=10 ;; b|B) v=11 ;; c|C) v=12 ;; d|D) v=13 ;; + e|E) v=14 ;; f|F) v=15 ;; + *) return 1 ;; + esac + + if [ $(( (v >> bit_in_char) & 1 )) -eq 1 ]; then + return 0 + else + return 1 + fi +} + +ev_contains() { + evstr="$1" + evbit="$2" + + evstr=$(print "$evstr" | tr -d ' \t\r\n') + + val=$(( 16#${evstr} )) + + if [ $(( (val >> evbit) & 1 )) -eq 1 ]; then + return 0 + else + return 1 + fi +} + +export "$@" + +SYSPATH="/sys/${DEVPATH}" +if [[ "$SYSPATH" = */event* ]]; then + SYSPATH="$SYSPATH/device" +fi + +EV=$(cat "$SYSPATH/capabilities/ev") +KEY=$(cat "$SYSPATH/capabilities/key") +ABS=$(cat "$SYSPATH/capabilities/abs") +REL=$(cat "$SYSPATH/capabilities/rel") +PROPS=$(cat "$SYSPATH/properties") + +case "$MODE" in + ID_INPUT) + echo 1 + ;; + ID_INPUT_MOUSE) + if bitmap_contains "$KEY" "$BTN_MOUSE" && ev_contains "$EV" "$EV_REL" ; then + echo 1 + else + exit 1 + fi + ;; + ID_INPUT_KEYBOARD) + for i in $ASCII_KEYS; do + if ! bitmap_contains "$KEY" "$i"; then + exit 1 + fi + done + echo 1 + ;; + ID_INPUT_TOUCHPAD) + if ev_contains "$EV" "$EV_REL" && bitmap_contains "$PROPS" "$INPUT_PROP_POINTER"; then + echo 1 + else + exit 1 + fi + ;; + ID_INPUT_TOUCHSCREEN) + if ev_contains "$EV" "$EV_ABS" && bitmap_contains "$PROPS" "$INPUT_PROP_DIRECT"; then + echo 1 + else + exit 1 + fi + ;; + ID_INPUT_JOYSTICK) + exit 1 + ;; + *) + exit 1 + ;; +esac