Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
22 changes: 16 additions & 6 deletions dissect/cstruct/lexer.py
Original file line number Diff line number Diff line change
Expand Up @@ -134,6 +134,14 @@ def tokenize(data: str) -> list[Token]:
return Lexer(data).tokenize()


def format_error_context(data: str, lineno: int, window: int = 1) -> str:
"""Format a snippet of the input data around a given line number for error messages."""
start = max(0, lineno - window - 2)
context = data.splitlines()[start:lineno]
no_len = len(str(lineno)) + 2
return "\n".join(f"{start + i + 1:>{no_len}}: {line}" for i, line in enumerate(context))


class Lexer:
"""Lexer compatible with C-like syntax for struct definitions and preprocessor directives."""

Expand Down Expand Up @@ -213,8 +221,10 @@ def _expect(self, *chars: str) -> str:

return self._take()

def _error(self, msg: str, *, line: int | None = None) -> LexerError:
return LexerError(f"line {line if line is not None else self._line}: {msg}")
def _error(self, msg: str, *, lineno: int | None = None) -> LexerError:
lineno = lineno if lineno is not None else self._line
content = format_error_context(self.data, lineno)
return LexerError(f"line {lineno}: {msg}\n{content}")

def _emit(self, type: TokenType, value: str, line: int, column: int = 0) -> None:
"""Emit a token with the given type and value at the specified line and column."""
Expand Down Expand Up @@ -402,15 +412,15 @@ def _read_preprocessor(self) -> None:
keyword = self._read_identifier()

if (token_type := _PP_KEYWORDS.get(keyword)) is None:
raise self._error(f"unknown preprocessor directive '#{keyword}'", line=line)
raise self._error(f"unknown preprocessor directive '#{keyword}'", lineno=line)

self._emit(token_type, keyword, line, col)

if token_type == TokenType.PP_DEFINE:
self._skip_whitespace_and_comments()

if not (name := self._read_identifier()):
raise self._error("expected identifier after '#define'", line=line)
raise self._error("expected identifier after '#define'", lineno=line)
self._emit(TokenType.IDENTIFIER, name, line)

self._skip_whitespace_and_comments()
Expand All @@ -434,7 +444,7 @@ def _read_preprocessor(self) -> None:
self._expect(">") # Consume closing `>`
value = f"<{value}>"
else:
raise self._error("expected include path after '#include'", line=line)
raise self._error("expected include path after '#include'", lineno=line)

self._emit(TokenType.STRING, value, line)

Expand Down Expand Up @@ -486,7 +496,7 @@ def tokenize(self) -> list[Token]:
self._emit(_SINGLE_CHARS[ch], self._take(), line, col)

else:
raise self._error(f"unexpected character {ch!r}", line=line)
raise self._error(f"unexpected character {ch!r}", lineno=line)

self._emit(TokenType.EOF, "", self._line, self._column)
return self._tokens
Expand Down
27 changes: 22 additions & 5 deletions dissect/cstruct/parser.py
Original file line number Diff line number Diff line change
Expand Up @@ -11,7 +11,7 @@
ParserError,
)
from dissect.cstruct.expression import Expression
from dissect.cstruct.lexer import IDENTIFIER_TYPES, TokenCursor, TokenType, tokenize
from dissect.cstruct.lexer import IDENTIFIER_TYPES, TokenCursor, TokenType, format_error_context, tokenize
from dissect.cstruct.types import BaseArray, BaseType, Field, Structure

if TYPE_CHECKING:
Expand Down Expand Up @@ -58,6 +58,7 @@ def __init__(self, cs: cstruct, compiled: bool = True, align: bool = False):
super().__init__(cs)
self.compiled = compiled
self.align = align
self._data = None

self._flags: list[str] = []
self._conditional_stack: list[tuple[Token, bool]] = []
Expand All @@ -74,6 +75,8 @@ def parse(self, data: str) -> None:

data = _join_line_continuations(data)

# Keep a reference for error messages
self._data = data
self._reset_tokens(tokenize(data))
self._parse()

Expand All @@ -88,7 +91,12 @@ def _at(self, *types: TokenType) -> bool:
return self._tokens[self._pos].type in types

def _error(self, msg: str, *, token: Token | None = None) -> ParserError:
return ParserError(f"line {(token if token is not None else self._tokens[self._pos]).line}: {msg}")
lineno = (token if token is not None else self._tokens[self._pos]).line
if self._data is None:
return ParserError(f"line {lineno}: {msg}")

content = format_error_context(self._data, lineno)
return ParserError(f"line {lineno}: {msg}\n{content}")

def _in_false_branch(self) -> bool:
"""Return whether we're currently in a false conditional branch."""
Expand Down Expand Up @@ -177,8 +185,9 @@ def _parse_define(self) -> None:
if value[-1] != quote:
raise self._error("unterminated bytes literal", token=token)

# Remove the leading b and surrounding quotes
value = ast.literal_eval(f"b{value[2:-1]!r}")
# Remove the leading b and surrounding quotes and flatten escape sequences
value = value[2:-1].encode().decode("unicode_escape")
value = ast.literal_eval(f"b{value!r}")
else:
try:
# Lazy mode, try to evaluate as a Python literal first (for simple constants)
Expand Down Expand Up @@ -385,7 +394,15 @@ def _parse_enum_or_flag(self) -> type[Enum | Flag]:
continue
self._assert_not_eof()

member_name = self._expect(TokenType.IDENTIFIER).value
# For historical reasons, we allow enum/flag member names to start with a digit
# E.g. `32BIT`
member_name = ""
if token := self._match(TokenType.NUMBER):
member_name += token.value
if token := self._match(TokenType.IDENTIFIER):
member_name += token.value
else:
member_name = self._expect(TokenType.IDENTIFIER).value

if self._match(TokenType.EQUALS):
expression = self._collect_until(TokenType.COMMA, TokenType.RBRACE)
Expand Down
18 changes: 18 additions & 0 deletions tests/test_lexer.py
Original file line number Diff line number Diff line change
@@ -1,5 +1,7 @@
from __future__ import annotations

import textwrap

import pytest

from dissect.cstruct.exception import LexerError
Expand Down Expand Up @@ -182,6 +184,22 @@ def test_error(src: str, match: str) -> None:
tokenize(src)


def test_error_context() -> None:
"""Test that a LexerError includes the surrounding lines as context."""
src = "struct test {\n uint32 a;\n @invalid;\n uint32 b;\n};"
with pytest.raises(LexerError, match="line 3: unexpected character '@'") as exc_info:
tokenize(src)

assert str(exc_info.value) == textwrap.dedent(
"""\
line 3: unexpected character '@'
1: struct test {
2: uint32 a;
3: @invalid;
""".rstrip()
)


def test_line_and_column_tracking() -> None:
"""Test that the lexer correctly tracks line and column numbers."""
src = "a\n b\nc"
Expand Down
55 changes: 53 additions & 2 deletions tests/test_parser.py
Original file line number Diff line number Diff line change
@@ -1,5 +1,7 @@
from __future__ import annotations

import textwrap

import pytest

from dissect.cstruct import cstruct
Expand Down Expand Up @@ -214,6 +216,27 @@ def test_typedef_enum(cs: cstruct) -> None:
assert cs.test_enum.VAL3 == 4


def test_enum_flag_digit_member_name(cs: cstruct) -> None:
# For historical reasons, we allow enum/flag member names to start with a digit, e.g. `32BIT`
cdef = """
enum test_enum : uint8 {
32BIT = 1,
64BIT = 2
};

flag test_flag : uint8 {
32BIT = 1,
64BIT = 2
};
"""
cs.load(cdef)

assert cs.test_enum["32BIT"] == 1
assert cs.test_enum["64BIT"] == 2
assert cs.test_flag["32BIT"] == 1
assert cs.test_flag["64BIT"] == 2


def test_define(cs: cstruct) -> None:
cdef = """
#define CONST 42
Expand All @@ -230,6 +253,7 @@ def test_define(cs: cstruct) -> None:
3)
#define QUOTES "\'\"a'b\""
#define ESCAPE "\\'\\"a'b\\"\\n"
#define BYTES_ESCAPE b"`\\n"
#define FUNC(x) ( x == 0 )
#define TERNARY(x) ( x ? 1 : 0 )
"""
Expand All @@ -247,6 +271,7 @@ def test_define(cs: cstruct) -> None:
assert cs.consts["MULTILINE"] == 6
assert cs.consts["QUOTES"] == "'\"a'b\""
assert cs.consts["ESCAPE"] == "'\"a'b\"\n"
assert cs.consts["BYTES_ESCAPE"] == b"`\n"
# We don't evaluate function-like macros yet, so they should be stored as their raw string representation
assert cs.consts["FUNC"] == "(x) ( x == 0 )"
assert cs.consts["TERNARY"] == "(x) ( x ? 1 : 0 )"
Expand Down Expand Up @@ -490,7 +515,7 @@ def test_preprocessor_in_struct_body(cs: cstruct) -> None:
assert cs.test.fields["bonus"].type == cs.uint64


def test_preprocessor_define_from_enum_in_struct() -> None:
def test_preprocessor_define_from_enum_in_struct(cs: cstruct) -> None:
"""Test #define referencing enum values used for conditional fields and array sizes."""
cdef = """
enum protocol : uint8 {
Expand Down Expand Up @@ -530,7 +555,6 @@ def test_preprocessor_define_from_enum_in_struct() -> None:
uint16 checksum;
};
"""
cs = cstruct()
cs.load(cdef)

assert cs.consts["PROTO"] == 6
Expand All @@ -544,3 +568,30 @@ def test_preprocessor_define_from_enum_in_struct() -> None:

assert cs.packet.fields["options"].type.num_entries == 4
assert cs.packet.fields["payload"].type.num_entries == 20


def test_error_context(cs: cstruct) -> None:
"""Test the context window around errors: 1 line before, the error line, up to 2 lines after."""
# Error on line 1: no preceding lines available; shows lines 1, 2, 3
src = "69\n#define A 1\n#define B 2\n#define C 3"
with pytest.raises(ParserError, match="line 1:") as exc_info:
cs.load(src)
assert str(exc_info.value) == textwrap.dedent(
"""\
line 1: unexpected token '69'
1: 69
""".rstrip()
)

# Error on line 3: shows 2 lines of preceding context plus the error line
src = "#define A 1\n#define B 2\n69\n#define C 3\n#define D 4\n#define E 5"
with pytest.raises(ParserError, match="line 3:") as exc_info:
cs.load(src)
assert str(exc_info.value) == textwrap.dedent(
"""\
line 3: unexpected token '69'
1: #define A 1
2: #define B 2
3: 69
""".rstrip()
)
Loading