diff --git a/dissect/cstruct/lexer.py b/dissect/cstruct/lexer.py index 977f6ed..30a03ec 100644 --- a/dissect/cstruct/lexer.py +++ b/dissect/cstruct/lexer.py @@ -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.""" @@ -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.""" @@ -402,7 +412,7 @@ 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) @@ -410,7 +420,7 @@ def _read_preprocessor(self) -> None: 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() @@ -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) @@ -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 diff --git a/dissect/cstruct/parser.py b/dissect/cstruct/parser.py index c9de97e..d6ef294 100644 --- a/dissect/cstruct/parser.py +++ b/dissect/cstruct/parser.py @@ -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: @@ -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]] = [] @@ -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() @@ -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.""" @@ -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) @@ -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) diff --git a/tests/test_lexer.py b/tests/test_lexer.py index de72e18..f08a584 100644 --- a/tests/test_lexer.py +++ b/tests/test_lexer.py @@ -1,5 +1,7 @@ from __future__ import annotations +import textwrap + import pytest from dissect.cstruct.exception import LexerError @@ -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" diff --git a/tests/test_parser.py b/tests/test_parser.py index 1f2f75c..9299e8a 100644 --- a/tests/test_parser.py +++ b/tests/test_parser.py @@ -1,5 +1,7 @@ from __future__ import annotations +import textwrap + import pytest from dissect.cstruct import cstruct @@ -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 @@ -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 ) """ @@ -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 )" @@ -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 { @@ -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 @@ -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() + )