use std::{
collections::HashMap,
env, fs,
io::{self, Cursor, Read},
path::PathBuf,
pin::Pin,
sync::{Arc, Mutex},
};
use crate::driver::{
AllList, DeleteItem, DeviceList, Empty, GetAll, GetItem, InsertDb, Item, ReadMsg, ResRead,
ResSend, SendMsg, SetLight, SystemInfo, VenderMsg, Version, WirelessLoopStatus,
driver_grpc_server::DriverGrpc,
};
use futures_core::Stream;
use tokio_stream as stream;
use tonic::{Request, Response, Status};
use crate::{
catalog::DeviceCatalog,
transport::{DeniedDeviceIo, DeviceIo},
};
type ResponseStream<T> = Pin<Box<dyn Stream<Item = Result<T, Status>> + Send>>;
type Database = HashMap<Vec<u8>, Vec<u8>>;
type Storage = HashMap<String, Database>;
pub struct DriverService {
catalog: Arc<dyn DeviceCatalog>,
device_io: Arc<dyn DeviceIo>,
storage: Mutex<Storage>,
storage_path: Option<PathBuf>,
}
impl DriverService {
pub fn new(catalog: Arc<dyn DeviceCatalog>) -> Self {
Self::with_device_io(catalog, Arc::new(DeniedDeviceIo))
}
pub fn with_device_io(catalog: Arc<dyn DeviceCatalog>, device_io: Arc<dyn DeviceIo>) -> Self {
let storage_path = storage_path();
let storage = storage_path
.as_ref()
.map(load_storage)
.transpose()
.unwrap_or_else(|error| {
eprintln!("database load failed: {error}");
None
})
.unwrap_or_default();
Self {
catalog,
device_io,
storage: Mutex::new(storage),
storage_path,
}
}
fn save_storage(&self, storage: &Storage) -> Result<(), Status> {
let Some(path) = &self.storage_path else {
return Ok(());
};
save_storage(path, storage).map_err(|error| Status::internal(error.to_string()))
}
}
#[tonic::async_trait]
#[allow(non_camel_case_types)]
impl DriverGrpc for DriverService {
type watchDevListStream = ResponseStream<DeviceList>;
type watchVenderStream = ResponseStream<VenderMsg>;
type watchSystemInfoStream = ResponseStream<SystemInfo>;
async fn watch_dev_list(
&self,
_request: Request<Empty>,
) -> Result<Response<Self::watchDevListStream>, Status> {
let snapshot = self.catalog.initial_snapshot();
Ok(Response::new(Box::pin(stream::iter([Ok(snapshot)]))))
}
async fn watch_vender(
&self,
_request: Request<Empty>,
) -> Result<Response<Self::watchVenderStream>, Status> {
let device_io = Arc::clone(&self.device_io);
let (sender, receiver) = tokio::sync::mpsc::channel(16);
tokio::task::spawn_blocking(move || {
loop {
match device_io.read_event(25) {
Ok(Some(msg)) => {
let prefix = msg
.iter()
.take(8)
.map(|byte| format!("{byte:02x}"))
.collect::<Vec<_>>()
.join(" ");
eprintln!("HID event length={} data={}", msg.len(), prefix);
if sender.blocking_send(Ok(VenderMsg { msg })).is_err() {
break;
}
}
Ok(None) => std::thread::sleep(std::time::Duration::from_millis(5)),
Err(error) => {
if sender
.blocking_send(Err(Status::unavailable(error)))
.is_err()
{
break;
}
break;
}
}
}
});
Ok(Response::new(Box::pin(
tokio_stream::wrappers::ReceiverStream::new(receiver),
)))
}
async fn watch_system_info(
&self,
_request: Request<Empty>,
) -> Result<Response<Self::watchSystemInfoStream>, Status> {
Ok(Response::new(Box::pin(stream::pending())))
}
async fn get_version(&self, _request: Request<Empty>) -> Result<Response<Version>, Status> {
Ok(Response::new(Version {
baseversion: env!("CARGO_PKG_VERSION").into(),
timestamp: "mock".into(),
}))
}
async fn send_msg(&self, request: Request<SendMsg>) -> Result<Response<ResSend>, Status> {
let request = request.into_inner();
let result =
self.device_io
.send(&request.device_path, &request.msg, request.check_sum_type);
let prefix = request
.msg
.iter()
.take(16)
.map(|byte| format!("{byte:02x}"))
.collect::<Vec<_>>()
.join(" ");
eprintln!(
"HID write checksum={} length={} data={} result={}",
request.check_sum_type,
request.msg.len(),
prefix,
result.as_ref().err().map_or("ok", String::as_str)
);
let err = result.err().unwrap_or_default();
Ok(Response::new(ResSend { err }))
}
async fn read_msg(&self, request: Request<ReadMsg>) -> Result<Response<ResRead>, Status> {
let request = request.into_inner();
let response = match self.device_io.read(&request.device_path) {
Ok(msg) => {
let prefix = msg
.iter()
.take(16)
.map(|byte| format!("{byte:02x}"))
.collect::<Vec<_>>()
.join(" ");
eprintln!("HID read length={} data={} result=ok", msg.len(), prefix);
ResRead {
err: String::new(),
msg,
}
}
Err(err) => {
eprintln!("HID read result={err}");
ResRead {
err,
msg: Vec::new(),
}
}
};
Ok(Response::new(response))
}
async fn send_raw_feature(
&self,
request: Request<SendMsg>,
) -> Result<Response<ResSend>, Status> {
let inner = request.get_ref();
eprintln!(
"sendRawFeature checksum={} command={:02x} length={}",
inner.check_sum_type,
inner.msg.first().copied().unwrap_or_default(),
inner.msg.len()
);
self.send_msg(request).await
}
async fn read_raw_feature(
&self,
request: Request<ReadMsg>,
) -> Result<Response<ResRead>, Status> {
eprintln!("readRawFeature path={}", request.get_ref().device_path);
self.read_msg(request).await
}
async fn set_light_type(&self, request: Request<SetLight>) -> Result<Response<Empty>, Status> {
let request = request.into_inner();
eprintln!(
"setLightType type={} screen={} path={}",
request.light_type, request.screen_id, request.device_path
);
if !matches!(request.light_type, 1 | 2) {
return Err(Status::permission_denied(
"only screen and idle helper-side modes are allowed",
));
}
Ok(Response::new(Empty {}))
}
async fn insert_db(&self, request: Request<InsertDb>) -> Result<Response<ResSend>, Status> {
let request = request.into_inner();
let mut storage = self
.storage
.lock()
.map_err(|_| Status::internal("database lock is poisoned"))?;
storage
.entry(request.dbpath)
.or_default()
.insert(request.key, request.value);
self.save_storage(&storage)?;
Ok(Response::new(ResSend { err: String::new() }))
}
async fn delete_item_from_db(
&self,
request: Request<DeleteItem>,
) -> Result<Response<ResSend>, Status> {
let request = request.into_inner();
let mut storage = self
.storage
.lock()
.map_err(|_| Status::internal("database lock is poisoned"))?;
if let Some(database) = storage.get_mut(&request.dbpath) {
database.remove(&request.key);
}
self.save_storage(&storage)?;
Ok(Response::new(ResSend { err: String::new() }))
}
async fn get_item_from_db(&self, request: Request<GetItem>) -> Result<Response<Item>, Status> {
let request = request.into_inner();
let storage = self
.storage
.lock()
.map_err(|_| Status::internal("database lock is poisoned"))?;
let value = storage
.get(&request.dbpath)
.and_then(|database| database.get(&request.key))
.cloned();
Ok(Response::new(match value {
Some(value) => Item {
value,
err_str: String::new(),
},
None => Item {
value: Vec::new(),
err_str: "not found".into(),
},
}))
}
async fn get_all_keys_from_db(
&self,
request: Request<GetAll>,
) -> Result<Response<AllList>, Status> {
let request = request.into_inner();
let storage = self
.storage
.lock()
.map_err(|_| Status::internal("database lock is poisoned"))?;
let data = storage
.get(&request.dbpath)
.map(|database| database.keys().cloned().collect())
.unwrap_or_default();
Ok(Response::new(AllList {
data,
err_str: String::new(),
}))
}
async fn get_all_values_from_db(
&self,
request: Request<GetAll>,
) -> Result<Response<AllList>, Status> {
let request = request.into_inner();
let storage = self
.storage
.lock()
.map_err(|_| Status::internal("database lock is poisoned"))?;
let data = storage
.get(&request.dbpath)
.map(|database| database.values().cloned().collect())
.unwrap_or_default();
Ok(Response::new(AllList {
data,
err_str: String::new(),
}))
}
async fn change_wireless_loop_status(
&self,
_request: Request<WirelessLoopStatus>,
) -> Result<Response<Empty>, Status> {
Ok(Response::new(Empty {}))
}
async fn clean_dev(&self, request: Request<ReadMsg>) -> Result<Response<ResSend>, Status> {
let request = request.into_inner();
let err = self
.device_io
.clean(&request.device_path)
.err()
.unwrap_or_default();
Ok(Response::new(ResSend { err }))
}
}
fn storage_path() -> Option<PathBuf> {
env::var_os("XDG_DATA_HOME")
.map(PathBuf::from)
.or_else(|| {
env::var_os("HOME").map(|home| PathBuf::from(home).join(".local").join("share"))
})
.map(|path| path.join("keydous-linux").join("web-driver.db"))
}
fn load_storage(path: &PathBuf) -> io::Result<Storage> {
let bytes = match fs::read(path) {
Ok(bytes) => bytes,
Err(error) if error.kind() == io::ErrorKind::NotFound => return Ok(HashMap::new()),
Err(error) => return Err(error),
};
let mut input = Cursor::new(bytes);
let mut magic = [0_u8; 8];
input.read_exact(&mut magic)?;
if &magic != b"KEYDOUS1" {
return Err(io::Error::new(
io::ErrorKind::InvalidData,
"invalid database header",
));
}
let entries = read_u32(&mut input)?;
let mut storage = HashMap::new();
for _ in 0..entries {
let dbpath = String::from_utf8(read_bytes(&mut input)?)
.map_err(|error| io::Error::new(io::ErrorKind::InvalidData, error))?;
let key = read_bytes(&mut input)?;
let value = read_bytes(&mut input)?;
storage
.entry(dbpath)
.or_insert_with(HashMap::new)
.insert(key, value);
}
if input.position() != input.get_ref().len() as u64 {
return Err(io::Error::new(
io::ErrorKind::InvalidData,
"trailing database data",
));
}
Ok(storage)
}
fn save_storage(path: &PathBuf, storage: &Storage) -> io::Result<()> {
let parent = path.parent().ok_or_else(|| {
io::Error::new(io::ErrorKind::InvalidInput, "database path has no parent")
})?;
fs::create_dir_all(parent)?;
let mut output = Vec::new();
output.extend_from_slice(b"KEYDOUS1");
let entries = storage.values().try_fold(0_u32, |count, database| {
let size = u32::try_from(database.len())
.map_err(|_| io::Error::new(io::ErrorKind::InvalidData, "too many database entries"))?;
count
.checked_add(size)
.ok_or_else(|| io::Error::new(io::ErrorKind::InvalidData, "too many database entries"))
})?;
output.extend_from_slice(&entries.to_le_bytes());
for (dbpath, database) in storage {
for (key, value) in database {
write_bytes(&mut output, dbpath.as_bytes())?;
write_bytes(&mut output, key)?;
write_bytes(&mut output, value)?;
}
}
let temporary = path.with_extension("db.tmp");
fs::write(&temporary, output)?;
fs::rename(temporary, path)
}
fn read_u32(input: &mut Cursor<Vec<u8>>) -> io::Result<u32> {
let mut bytes = [0_u8; 4];
input.read_exact(&mut bytes)?;
Ok(u32::from_le_bytes(bytes))
}
fn read_bytes(input: &mut Cursor<Vec<u8>>) -> io::Result<Vec<u8>> {
let length = read_u32(input)? as usize;
if length > 64 * 1024 * 1024 {
return Err(io::Error::new(
io::ErrorKind::InvalidData,
"database field is too large",
));
}
let mut bytes = vec![0_u8; length];
input.read_exact(&mut bytes)?;
Ok(bytes)
}
fn write_bytes(output: &mut Vec<u8>, bytes: &[u8]) -> io::Result<()> {
let length = u32::try_from(bytes.len())
.map_err(|_| io::Error::new(io::ErrorKind::InvalidData, "database field is too large"))?;
output.extend_from_slice(&length.to_le_bytes());
output.extend_from_slice(bytes);
Ok(())
}
#[cfg(test)]
mod tests {
use std::{collections::HashMap, sync::Arc};
use crate::driver::{Empty, dj_dev::Oneofdev, driver_grpc_server::DriverGrpc};
use tempfile::tempdir;
use tokio_stream::StreamExt;
use tonic::Request;
use crate::{
catalog::SimulatedCatalog,
service::{DriverService, load_storage, save_storage},
};
#[tokio::test]
async fn watch_dev_list_starts_with_initial_snapshot() {
let service = DriverService::new(Arc::new(SimulatedCatalog));
let response = service
.watch_dev_list(Request::new(Empty {}))
.await
.unwrap();
let snapshot = response.into_inner().next().await.unwrap().unwrap();
let Some(Oneofdev::Dev(device)) = &snapshot.devlist[0].oneofdev else {
panic!("expected direct device");
};
assert_eq!((device.vid, device.pid), (0x3151, 0x5030));
}
#[test]
fn web_driver_database_survives_restart() {
let directory = tempdir().unwrap();
let path = directory.path().join("web-driver.db");
let storage = HashMap::from([(
"/driver/CONFIG".into(),
HashMap::from([(b"NJ98-CP V4".to_vec(), br#"{"profile":3}"#.to_vec())]),
)]);
save_storage(&path, &storage).unwrap();
assert_eq!(load_storage(&path).unwrap(), storage);
}
}