use std::str::FromStr;
use crate::node::{CharGroup, Key, Modifier, Node, KEY_SEP};
type ParserFn<T> = fn(&mut Parser) -> Result<Option<T>, ParseError>;
#[derive(Debug, PartialEq, Clone)]
pub struct ParseError {
pub message: String,
pub position: usize,
}
impl std::fmt::Display for ParseError {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
write!(
f,
"Parse error at position {}: {}",
self.position, self.message
)
}
}
impl std::error::Error for ParseError {}
struct Parser<'a> {
input: &'a str,
position: usize,
}
impl<'a> Parser<'a> {
pub fn new(input: &'a str) -> Self {
Self { input, position: 0 }
}
pub fn peek(&self) -> Option<char> {
self.input.chars().next()
}
pub fn peek_at(&self, n: usize) -> Option<char> {
self.input.chars().nth(n)
}
pub fn next(&mut self) -> Option<char> {
if let Some(ch) = self.peek() {
self.position += ch.len_utf8();
self.input = &self.input[ch.len_utf8()..];
Some(ch)
} else {
None
}
}
pub fn is_end(&self) -> bool {
self.input.is_empty()
}
pub fn take(&mut self, expected: char) -> Result<(), ParseError> {
match self.next() {
Some(ch) if ch == expected => Ok(()),
Some(ch) => Err(ParseError {
message: format!("expected '{expected}', found '{ch}'"),
position: self.position - ch.len_utf8(),
}),
None => Err(ParseError {
message: format!("expected '{expected}', found end of input"),
position: self.position,
}),
}
}
pub fn try_parse<T, F>(&mut self, f: F) -> Result<Option<T>, ParseError>
where
F: FnOnce(&mut Parser<'a>) -> Result<Option<T>, ParseError>,
{
let snapshot = (self.input, self.position);
match f(self) {
Ok(Some(val)) => Ok(Some(val)),
Ok(None) | Err(_) => {
self.input = snapshot.0;
self.position = snapshot.1;
Ok(None)
}
}
}
pub fn take_while<F>(&mut self, predicate: F) -> String
where
F: Fn(char) -> bool,
{
let mut result = String::new();
while let Some(ch) = self.peek() {
if predicate(ch) {
result.push(ch);
self.next();
} else {
break;
}
}
result
}
pub fn alt<T>(&mut self, parsers: &[ParserFn<T>]) -> Result<Option<T>, ParseError> {
for p in parsers {
match p(self)? {
Some(value) => return Ok(Some(value)),
None => continue,
}
}
Ok(None)
}
pub fn error(&self, message: String) -> ParseError {
ParseError {
message,
position: self.position,
}
}
}
pub fn parse(s: &str) -> Result<Node, ParseError> {
let mut parser = Parser::new(s);
let node = parse_node(&mut parser)?;
if !parser.is_end() {
return Err(parser.error(format!(
"expect end of input, found: {}",
parser.peek().unwrap()
)));
}
Ok(node)
}
fn parse_node(parser: &mut Parser) -> Result<Node, ParseError> {
let mut modifiers = 0u8;
for _ in 0..4 {
if let Some(modifier) = try_parse_modifier(parser)? {
modifiers |= modifier as u8;
} else {
break;
}
}
let key = parse_key(parser)?;
Ok(Node::new(modifiers, key))
}
fn try_parse_modifier(parser: &mut Parser) -> Result<Option<Modifier>, ParseError> {
parser.try_parse(|p| {
let name = p.take_while(|ch| ch.is_ascii_alphabetic());
let Ok(modifier) = name.parse::<Modifier>() else {
return Ok(None);
};
p.take(KEY_SEP)?;
Ok(Some(modifier))
})
}
fn parse_key(parser: &mut Parser) -> Result<Key, ParseError> {
match parser.alt(&[
try_parse_fn_key,
try_parse_named_key,
try_parse_group,
try_parse_char,
])? {
Some(key) => Ok(key),
None => Err(parser.error("expected a valid key".to_string())),
}
}
fn try_parse_fn_key(parser: &mut Parser) -> Result<Option<Key>, ParseError> {
if parser.peek() != Some('f') || parser.peek_at(1).is_none() {
return Ok(None);
}
parser.take('f')?;
parser.try_parse(|p| {
let num = p.take_while(|ch| ch.is_ascii_digit());
match num.parse::<u8>() {
Ok(n) if n <= 12 => Ok(Some(Key::F(n))),
_ => Err(p.error("invalid function key number (must be 0-12)".to_string())),
}
})
}
fn try_parse_named_key(parser: &mut Parser) -> Result<Option<Key>, ParseError> {
parser.try_parse(|p| {
let name = p.take_while(|ch| ch.is_ascii_alphabetic());
if name.len() < 2 {
return Ok(None);
}
match name.parse::<Key>() {
Ok(key) => Ok(Some(key)),
Err(_) => Ok(None),
}
})
}
fn try_parse_group(parser: &mut Parser) -> Result<Option<Key>, ParseError> {
if parser.peek() != Some('@') || parser.peek_at(1).is_none() {
return Ok(None);
}
parser.take('@')?;
let group_name = parser.take_while(|ch| ch.is_ascii_alphabetic());
let group = match group_name.parse::<CharGroup>() {
Ok(group) => Key::Group(group),
Err(_) => return Err(parser.error(format!("unknown char group: '@{group_name}'"))),
};
Ok(Some(group))
}
fn try_parse_char(parser: &mut Parser) -> Result<Option<Key>, ParseError> {
if let Some(ch) = parser.peek() {
if ch.is_ascii() {
parser.next();
Ok(Some(Key::Char(ch)))
} else {
Ok(None)
}
} else {
Ok(None)
}
}
pub fn parse_seq(s: &str) -> Result<Vec<Node>, ParseError> {
str::split_whitespace(s).map(parse).collect()
}
impl FromStr for Node {
type Err = ParseError;
fn from_str(s: &str) -> Result<Self, Self::Err> {
parse(s)
}
}
#[cfg(test)]
mod tests {
use serde::Deserialize;
use crate::parser::{CharGroup, Key, Modifier, Node};
use super::{parse, ParseError};
#[test]
fn test_parse() {
let err = |message: &str, position: usize| {
Err::<Node, ParseError>(ParseError {
message: message.to_string(),
position,
})
};
[
("alt-f", Ok(Node::new(Modifier::Alt as u8, Key::Char('f')))),
("space", Ok(Node::new(0, Key::Space))),
("delta", err("expect end of input, found: e", 1)),
(
"shift-a",
Ok(Node::new(Modifier::Shift as u8, Key::Char('a'))),
),
("shift-a-delete", err("expect end of input, found: -", 7)),
("al", err("expect end of input, found: l", 1)),
]
.iter()
.for_each(|(input, result)| {
let output = parse(input);
assert_eq!(result, &output);
});
}
#[test]
fn test_parse_seq() {
[
("ctrl-b", Ok(vec![parse("ctrl-b").unwrap()])),
(
"ctrl-b l",
Ok(vec![parse("ctrl-b").unwrap(), parse("l").unwrap()]),
),
("ctrl-b -l", Err(parse("-l").unwrap_err())), ("b", Ok(vec![parse("b").unwrap()])),
]
.iter()
.for_each(|(s, v)| assert_eq!(&super::parse_seq(s), v));
}
#[test]
fn test_parse_fn_key() {
(0..=12).for_each(|n| {
let input = format!("f{n}");
let result = parse(&input);
assert_eq!(Key::F(n), result.unwrap().key);
});
[13, 15].iter().for_each(|n| {
let input = format!("f{n}");
let result = parse(&input);
assert!(result.is_err());
});
}
#[test]
fn test_parse_enum() {
[("up", Key::Up), ("esc", Key::Esc), ("del", Key::Delete)]
.iter()
.for_each(|(s, key)| {
let result = parse(s);
assert_eq!(&result.unwrap().key, key);
});
}
#[test]
fn test_parse_char_groups() {
[
("@digit", Key::Group(CharGroup::Digit)),
("@lower", Key::Group(CharGroup::Lower)),
("@upper", Key::Group(CharGroup::Upper)),
("@alpha", Key::Group(CharGroup::Alpha)),
("@alnum", Key::Group(CharGroup::Alnum)),
("@any", Key::Group(CharGroup::Any)),
]
.iter()
.for_each(|(input, expected_key)| {
let result = parse(input);
assert_eq!(&result.unwrap().key, expected_key);
});
let result = parse("@invalid");
assert!(result.is_err());
assert!(result
.unwrap_err()
.message
.contains("unknown char group: '@invalid'"));
let result = parse("@x");
assert!(result.is_err());
assert!(result
.unwrap_err()
.message
.contains("unknown char group: '@x'"));
}
#[test]
fn test_format() {
[
(Node::new(0, Key::F(3)), "f3"),
(Node::new(0, Key::Delete), "delete"),
(Node::new(0, Key::Space), "space"),
(Node::new(0, Key::Char('g')), "g"),
(Node::new(0, Key::Char('#')), "#"),
(Node::new(0, Key::Group(CharGroup::Digit)), "@digit"),
(Node::new(0, Key::Group(CharGroup::Lower)), "@lower"),
(Node::new(Modifier::Alt as u8, Key::Char('f')), "alt-f"),
(
Node::new(Modifier::Alt as u8, Key::Group(CharGroup::Alpha)),
"alt-@alpha",
),
(
Node::new(Modifier::Shift as u8 | Modifier::Cmd as u8, Key::Char('f')),
"cmd-shift-f",
),
]
.iter()
.for_each(|(node, expected)| {
assert_eq!(expected, &format!("{node}"));
});
}
#[test]
fn test_deserialize() {
use std::collections::HashMap;
#[derive(Deserialize, Debug)]
struct Test {
keys: HashMap<Node, String>,
}
let result: Test = toml::from_str(
r#"
[keys]
alt-d = "a"
cmd-shift-del = "b"
shift-cmd-del = "b" # this is the same as previous one
delete = "d"
"@digit" = "number"
"alt-@lower" = "alt-lowercase"
"#,
)
.unwrap();
[
Node::new(Modifier::Alt as u8, Key::Char('d')),
Node::new(Modifier::Cmd as u8 | Modifier::Shift as u8, Key::Delete),
Node::new(0, Key::Delete),
Node::new(0, Key::Group(CharGroup::Digit)),
Node::new(Modifier::Alt as u8, Key::Group(CharGroup::Lower)),
]
.iter()
.for_each(|n| {
let (key, _) = result.keys.get_key_value(n).unwrap();
assert_eq!(key, n);
});
}
#[test]
fn test_parse_str() {
[
(Node::new(0, Key::F(3)), "f3"),
(Node::new(0, Key::Delete), "delete"),
(Node::new(0, Key::Space), "space"),
]
.iter()
.for_each(|(expected, input)| {
let node = input.parse::<Node>().unwrap();
assert_eq!(expected, &node);
});
}
}