use crate::backend::MemoryOps;
use crate::error::{Error, Result};
use crate::symbols::{ParsedType, SymbolStore};
use crate::target::Target;
use crate::types::{Dtb, VirtAddr};
use crate::ui;
use owo_colors::OwoColorize;
use std::ops::Range;
use winnow::Parser;
use winnow::combinator::{alt, peek, repeat};
use winnow::error::{ErrMode, ModalResult, ParserError};
use winnow::stream::{LocatingSlice, Location, Stream};
use winnow::token::{literal, one_of, take_till, take_while};
#[derive(Debug, Clone, PartialEq)]
pub enum ExprType {
Byte,
Word,
Dword,
Qword,
Struct(String),
Pointer(Box<ExprType>),
}
#[derive(Debug, Clone, PartialEq)]
pub enum Expr {
Literal(VirtAddr),
Symbol(String),
Register(String),
Deref(Box<Expr>),
FieldAccess(Box<Expr>, String),
Index(Box<Expr>, u64),
Add(Box<Expr>, Box<Expr>),
Sub(Box<Expr>, Box<Expr>),
Cast(Box<Expr>, ExprType),
}
#[derive(Debug, Clone, PartialEq, Eq)]
pub struct ExprParseError {
pub span: Range<usize>,
pub label: String,
}
impl ExprParseError {
fn new(span: Range<usize>, label: impl Into<String>) -> Self {
Self {
span,
label: label.into(),
}
}
pub fn render(&self, input: &str) -> String {
let start = self.span.start.min(input.len());
let end = self.span.end.min(input.len()).max(start + 1);
let caret_width = end.saturating_sub(start).max(1);
format!(
"{}\n{}{} {}",
input,
" ".repeat(start),
"^".repeat(caret_width).red(),
ui::muted(&self.label)
)
}
}
impl Expr {
pub fn eval(input: &str, context: &Target) -> Result<VirtAddr> {
Self::parse(input)?.resolve(context)
}
pub fn parse(input: &str) -> Result<Self> {
Self::parse_detailed(input).map_err(|err| Error::InvalidExpression(err.render(input)))
}
pub fn parse_detailed(input: &str) -> std::result::Result<Self, ExprParseError> {
let mut input = LocatingSlice::new(input);
let expr = parse_additive
.parse_next(&mut input)
.map_err(unwrap_parse_error)?;
ws0.parse_next(&mut input).map_err(unwrap_parse_error)?;
if input.peek_token().is_some() {
return Err(error_at(&input, "expected end of expression"));
}
Ok(expr)
}
fn is_type_name(s: &str) -> bool {
let s = s.trim().trim_end_matches('*').trim();
!s.is_empty()
&& s.starts_with(|c: char| c.is_ascii_alphabetic() || c == '_')
&& s.chars().all(|c| c.is_ascii_alphanumeric() || c == '_')
}
fn parse_type(type_str: &str) -> Result<ExprType> {
let type_str = type_str.trim();
if let Some(stripped) = type_str.strip_suffix('*') {
let inner_type = Self::parse_type(stripped)?;
return Ok(ExprType::Pointer(Box::new(inner_type)));
}
match type_str.to_lowercase().as_str() {
"byte" | "u8" | "uchar" | "char" | "boolean" | "uint8_t" | "int8_t" => {
Ok(ExprType::Byte)
}
"word" | "u16" | "ushort" | "short" | "wchar" | "uint16_t" | "int16_t" => {
Ok(ExprType::Word)
}
"dword" | "u32" | "ulong" | "long" | "uint" | "int" | "uint32_t" | "int32_t" => {
Ok(ExprType::Dword)
}
"qword" | "u64" | "dword64" | "ulong64" | "longlong" | "ulonglong" | "pvoid"
| "size_t" | "uint64_t" | "int64_t" | "usize" => Ok(ExprType::Qword),
_ => Ok(ExprType::Struct(type_str.to_string())),
}
}
pub fn resolve(&self, context: &Target) -> Result<VirtAddr> {
match self {
Expr::Literal(addr) => Ok(*addr),
Expr::Symbol(name) => {
if let Some(addr) = context
.symbols
.find_symbol_across_modules(context.current_dtb(), name)
{
return Ok(addr);
}
if let Some(value) = Self::parse_bare_hex_literal(name) {
return Ok(VirtAddr(value));
}
if let Some(value) = context.registers.as_ref().and_then(|r| r.get(name)) {
return Ok(VirtAddr(*value));
}
if let Some(value) = context.builtin_variable_value(name) {
return Ok(VirtAddr(value));
}
Err(Error::SymbolNotFound(name.clone()))
}
Expr::Register(name) => {
if !name.is_empty() && name.chars().all(|c| c.is_ascii_digit()) {
let idx: usize = name.parse().map_err(|_| {
Error::InvalidExpression(format!("invalid result index ${name}"))
})?;
return context
.results
.get(idx)
.copied()
.map(VirtAddr)
.ok_or_else(|| {
Error::InvalidExpression(format!(
"no result ${name} ({} available)",
context.results.len()
))
});
}
if let Some(value) = context.registers.as_ref().and_then(|r| r.get(name)) {
return Ok(VirtAddr(*value));
}
if let Some(var) = context.user_vars.get(name) {
return Ok(VirtAddr(var.value));
}
if let Some(value) = context.builtin_variable_value(name) {
return Ok(VirtAddr(value));
}
if context.registers.is_none() {
return Err(Error::InvalidExpression(
"registers unavailable while VM is running".into(),
));
}
Err(Error::RegisterNotFound(name.clone()))
}
Expr::Deref(inner) => {
let addr = inner.resolve(context)?;
let mem = context.current_process().memory();
let val: u64 = mem.read(addr)?;
Ok(VirtAddr(val))
}
Expr::FieldAccess(base, field_name) => {
let base_addr = base.resolve(context)?;
let base_type_name = base
.resolve_type(&context.symbols, context.current_dtb())
.ok_or_else(|| {
Error::InvalidExpression(
"field access requires explicit cast: e.g., (TYPE)expr->field".into(),
)
})?;
let type_info = context
.symbols
.find_type_across_modules(context.current_dtb(), &base_type_name)
.ok_or_else(|| Error::StructNotFound(base_type_name))?;
let offset = type_info
.field_offset(field_name)
.map_err(|_| Error::FieldNotFound(field_name.clone()))?;
Ok(base_addr + offset)
}
Expr::Index(base, index) => {
let base_addr = base.resolve(context)?;
let elem_size = Self::resolve_element_size(base, context);
Ok(base_addr + index * elem_size)
}
Expr::Add(lhs, rhs) => {
let base = lhs.resolve(context)?;
let value = rhs.resolve(context)?;
Ok(base + value.0)
}
Expr::Sub(lhs, rhs) => {
let base = lhs.resolve(context)?;
let value = rhs.resolve(context)?;
Ok(base - value.0)
}
Expr::Cast(expr, _) => expr.resolve(context),
}
}
pub fn resolve_type(&self, symbols: &SymbolStore, dtb: Dtb) -> Option<String> {
match self {
Expr::Cast(_, expr_type) => Some(Self::expr_type_to_name(expr_type)),
Expr::FieldAccess(base, field_name) => {
let base_type_name = base.resolve_type(symbols, dtb)?;
let type_info = symbols.find_type_across_modules(dtb, &base_type_name)?;
let field_info = type_info.fields.get(field_name)?;
Self::struct_type_name(&field_info.type_data)
}
_ => None,
}
}
pub fn complete_fields(&self, symbols: &SymbolStore, dtb: Dtb, prefix: &str) -> Vec<String> {
let type_name = match self.resolve_type(symbols, dtb) {
Some(name) => name,
None => return vec![],
};
let type_info = match symbols.find_type_across_modules(dtb, &type_name) {
Some(info) => info,
None => return vec![],
};
let mut fields: Vec<String> = type_info
.fields
.keys()
.filter(|f| f.starts_with(prefix))
.cloned()
.collect();
fields.sort();
fields
}
fn resolve_element_size(expr: &Expr, context: &Target) -> u64 {
match expr {
Expr::Cast(_, expr_type) => Self::expr_type_size(expr_type, context),
_ => 1,
}
}
fn expr_type_size(expr_type: &ExprType, context: &Target) -> u64 {
match expr_type {
ExprType::Byte => 1,
ExprType::Word => 2,
ExprType::Dword => 4,
ExprType::Qword => 8,
ExprType::Pointer(_) => 8,
ExprType::Struct(name) => {
let lookup = if name.starts_with('_') {
name.clone()
} else {
format!("_{name}")
};
context
.symbols
.find_type_across_modules(context.current_dtb(), &lookup)
.map(|t| t.size as u64)
.unwrap_or(1)
}
}
}
fn expr_type_to_name(expr_type: &ExprType) -> String {
match expr_type {
ExprType::Byte => "byte".to_string(),
ExprType::Word => "word".to_string(),
ExprType::Dword => "dword".to_string(),
ExprType::Qword => "qword".to_string(),
ExprType::Struct(name) => {
if name.starts_with('_') {
name.clone()
} else {
format!("_{name}")
}
}
ExprType::Pointer(inner) => Self::expr_type_to_name(inner),
}
}
fn struct_type_name(type_data: &ParsedType) -> Option<String> {
match type_data {
ParsedType::Struct(name) | ParsedType::Union(name) => Some(name.clone()),
ParsedType::Pointer(inner) => Self::struct_type_name(inner),
ParsedType::Array(inner, _) => Self::struct_type_name(inner),
_ => None,
}
}
fn parse_bare_hex_literal(s: &str) -> Option<u64> {
let s = s.trim();
let has_hex_letter = s
.chars()
.any(|ch| ch.is_ascii_hexdigit() && ch.is_ascii_alphabetic());
if !has_hex_letter || !s.chars().all(|ch| ch.is_ascii_hexdigit()) {
return None;
}
u64::from_str_radix(s, 16).ok()
}
}
type ExprInput<'a> = LocatingSlice<&'a str>;
type ParseResult<T> = ModalResult<T, ExprParseError>;
impl<'a> ParserError<ExprInput<'a>> for ExprParseError {
type Inner = Self;
fn from_input(input: &ExprInput<'a>) -> Self {
error_at(input, "expected expression")
}
fn into_inner(self) -> std::result::Result<Self::Inner, Self> {
Ok(self)
}
fn or(self, other: Self) -> Self {
if other.span.start >= self.span.start {
other
} else {
self
}
}
}
enum Suffix {
Field(String),
Index(u64),
}
fn parse_additive(input: &mut ExprInput<'_>) -> ParseResult<Expr> {
let expr = parse_postfix.parse_next(input)?;
let terms: Vec<(char, Expr)> = repeat(0.., parse_additive_tail).parse_next(input)?;
Ok(terms
.into_iter()
.fold(expr, |expr, (op, rhs)| match (op, rhs) {
('+', rhs) => Expr::Add(Box::new(expr), Box::new(rhs)),
('-', rhs) => Expr::Sub(Box::new(expr), Box::new(rhs)),
_ => unreachable!(),
}))
}
fn parse_additive_tail(input: &mut ExprInput<'_>) -> ParseResult<(char, Expr)> {
ws0.parse_next(input)?;
let op = one_of(['+', '-']).parse_next(input)?;
let rhs = parse_postfix.parse_next(input).map_err(ErrMode::cut)?;
Ok((op, rhs))
}
fn parse_postfix(input: &mut ExprInput<'_>) -> ParseResult<Expr> {
let expr = parse_prefix.parse_next(input)?;
let suffixes: Vec<Suffix> =
repeat(0.., alt((parse_field_suffix, parse_index_suffix))).parse_next(input)?;
Ok(suffixes
.into_iter()
.fold(expr, |expr, suffix| match suffix {
Suffix::Field(field) => Expr::FieldAccess(Box::new(expr), field),
Suffix::Index(index) => Expr::Index(Box::new(expr), index),
}))
}
fn parse_field_suffix(input: &mut ExprInput<'_>) -> ParseResult<Suffix> {
ws0.parse_next(input)?;
literal("->").parse_next(input)?;
ws0.parse_next(input)?;
let field = parse_field_name
.parse_next(input)
.map_err(|_| ErrMode::Cut(error_at(input, "expected field name after '->'")))?;
Ok(Suffix::Field(field.to_string()))
}
fn parse_index_suffix(input: &mut ExprInput<'_>) -> ParseResult<Suffix> {
ws0.parse_next(input)?;
one_of('[').parse_next(input)?;
ws0.parse_next(input)?;
let index = parse_number_literal(input, "expected numeric index")?;
ws0.parse_next(input)?;
expect_char(input, ']', "expected ']'")?;
Ok(Suffix::Index(index))
}
fn parse_prefix(input: &mut ExprInput<'_>) -> ParseResult<Expr> {
ws0.parse_next(input)?;
alt((parse_deref_prefix, parse_poi, parse_cast, parse_atom)).parse_next(input)
}
fn parse_deref_prefix(input: &mut ExprInput<'_>) -> ParseResult<Expr> {
one_of('*').parse_next(input)?;
let inner = parse_prefix.parse_next(input).map_err(ErrMode::cut)?;
Ok(Expr::Deref(Box::new(inner)))
}
fn parse_poi(input: &mut ExprInput<'_>) -> ParseResult<Expr> {
literal("poi").parse_next(input)?;
ws0.parse_next(input)?;
one_of('(').parse_next(input)?;
ws0.parse_next(input)?;
let inner = parse_additive.parse_next(input).map_err(ErrMode::cut)?;
ws0.parse_next(input)?;
expect_char(input, ')', "expected ')' after poi expression")?;
Ok(Expr::Deref(Box::new(inner)))
}
fn parse_cast(input: &mut ExprInput<'_>) -> ParseResult<Expr> {
let expr_type = parse_cast_type.parse_next(input)?;
ws0.parse_next(input)?;
peek(parse_operand_start).parse_next(input)?;
let base = parse_prefix.parse_next(input).map_err(ErrMode::cut)?;
Ok(Expr::Cast(Box::new(base), expr_type))
}
fn parse_cast_type(input: &mut ExprInput<'_>) -> ParseResult<ExprType> {
one_of('(').parse_next(input)?;
ws0.parse_next(input)?;
let span = take_till(1.., ')').parse_next(input)?;
let type_str = span.trim();
if !Expr::is_type_name(type_str) {
return Err(ErrMode::Backtrack(error_at(input, "expected expression")));
}
let expr_type = Expr::parse_type(type_str)
.map_err(|_| ErrMode::Cut(error_at(input, "invalid cast type")))?;
expect_char(input, ')', "expected ')'")?;
Ok(expr_type)
}
fn parse_atom(input: &mut ExprInput<'_>) -> ParseResult<Expr> {
alt((
parse_group,
parse_register,
parse_literal_expr,
parse_symbol_expr,
))
.parse_next(input)
}
fn parse_group(input: &mut ExprInput<'_>) -> ParseResult<Expr> {
one_of('(').parse_next(input)?;
ws0.parse_next(input)?;
let expr = parse_additive.parse_next(input).map_err(ErrMode::cut)?;
ws0.parse_next(input)?;
expect_char(input, ')', "expected ')'")?;
Ok(expr)
}
fn parse_register(input: &mut ExprInput<'_>) -> ParseResult<Expr> {
let (sigil, span) = one_of(['$', '@']).with_span().parse_next(input)?;
let name = parse_register_name.parse_next(input).map_err(|_| {
ErrMode::Cut(ExprParseError::new(
span,
format!("expected register name after '{sigil}'"),
))
})?;
Ok(Expr::Register(name.to_string()))
}
fn parse_literal_expr(input: &mut ExprInput<'_>) -> ParseResult<Expr> {
peek(one_of('0'..='9')).parse_next(input)?;
let value = parse_number_literal(input, "expected numeric literal").map_err(ErrMode::cut)?;
Ok(Expr::Literal(VirtAddr(value)))
}
fn parse_symbol_expr(input: &mut ExprInput<'_>) -> ParseResult<Expr> {
let symbol = parse_symbol_name
.parse_next(input)
.map_err(|_| ErrMode::Backtrack(error_at(input, "expected expression")))?;
Ok(Expr::Symbol(symbol.to_string()))
}
fn parse_number_literal(input: &mut ExprInput<'_>, label: &'static str) -> ParseResult<u64> {
let (token, span) = parse_number_token
.with_span()
.parse_next(input)
.map_err(|_| ErrMode::Backtrack(error_at(input, label)))?;
if token.starts_with("0x") || token.starts_with("0X") {
return u64::from_str_radix(&token[2..], 16)
.map_err(|_| ErrMode::Cut(ExprParseError::new(span, "invalid hex literal")));
}
if token.starts_with("0b") || token.starts_with("0B") {
return u64::from_str_radix(&token[2..], 2)
.map_err(|_| ErrMode::Cut(ExprParseError::new(span, "invalid binary literal")));
}
token
.parse::<u64>()
.map_err(|_| ErrMode::Cut(ExprParseError::new(span, "invalid numeric literal")))
}
fn parse_operand_start(input: &mut ExprInput<'_>) -> ParseResult<()> {
ws0.parse_next(input)?;
alt((
one_of(['(', '*', '$', '@']).void(),
one_of('0'..='9').void(),
take_while(1.., symbol_char).void(),
))
.parse_next(input)
}
fn parse_register_name<'a>(input: &mut ExprInput<'a>) -> ParseResult<&'a str> {
take_while(1.., |c: char| c.is_ascii_alphanumeric() || c == '_').parse_next(input)
}
fn parse_field_name<'a>(input: &mut ExprInput<'a>) -> ParseResult<&'a str> {
take_while(1.., |c: char| !is_expr_boundary(c) && c != ']').parse_next(input)
}
fn parse_symbol_name<'a>(input: &mut ExprInput<'a>) -> ParseResult<&'a str> {
take_while(1.., symbol_char).parse_next(input)
}
fn parse_number_token<'a>(input: &mut ExprInput<'a>) -> ParseResult<&'a str> {
take_while(1.., |c: char| !is_expr_boundary(c) && c != ']').parse_next(input)
}
fn ws0(input: &mut ExprInput<'_>) -> ParseResult<()> {
take_while(0.., char::is_whitespace)
.void()
.parse_next(input)
}
fn expect_char(input: &mut ExprInput<'_>, expected: char, label: &'static str) -> ParseResult<()> {
let parsed: ParseResult<char> = one_of(expected).parse_next(input);
parsed
.map(|_| ())
.map_err(|_| ErrMode::Cut(error_at(input, label)))
}
fn error_at(input: &ExprInput<'_>, label: impl Into<String>) -> ExprParseError {
let start = input.current_token_start();
let end = input
.peek_token()
.map(|ch| start + ch.len_utf8())
.unwrap_or(start + 1);
ExprParseError::new(start..end, label)
}
fn unwrap_parse_error(err: ErrMode<ExprParseError>) -> ExprParseError {
match err {
ErrMode::Backtrack(err) | ErrMode::Cut(err) => err,
ErrMode::Incomplete(_) => ExprParseError::new(0..1, "incomplete expression"),
}
}
fn symbol_char(ch: char) -> bool {
!is_expr_boundary(ch) && ch != '*' && ch != ']'
}
fn is_expr_boundary(ch: char) -> bool {
ch.is_whitespace() || matches!(ch, '(' | ')' | '[' | '+' | '-')
}
#[cfg(test)]
mod tests;