Coverage for src/renaissance/syntax_tree/ast_rewriter.py: 87%
307 statements
« prev ^ index » next coverage.py v7.16.1, created at 2026-10-02 13:12 +0000
« prev ^ index » next coverage.py v7.16.1, created at 2026-10-02 13:12 +0000
1"""AI: Rewriter that translates AST-level replace/remove/insert actions into byte-level edits."""
3import re
4import sys
5from collections.abc import Sequence
6from enum import Enum
7from typing import Protocol, Self, runtime_checkable
9from more_itertools import flatten
11from renaissance.common import Rewriter
12from renaissance.utils.text_utils import TextUtils
14from .match_finder import PatternMatch
15from .semantic_kind import SemanticKind
18@runtime_checkable
19class Rewritable(Protocol):
20 """AI: Structural protocol describing the offset/text shape required for byte-level rewriting."""
22 offset: int
23 end_offset: int
24 extended_end_offset: int
25 filename: str
26 parent: Self
27 text: str
30class _RewriteActionType(Enum):
31 REPLACE = 1
32 INSERT_BEFORE = 2
33 INSERT_AFTER = 3
34 REMOVE = 4 # TODO: Why needed? Why isn't a REMOVE Action Type just a REPLACE Action Type (with an empty string)?
37DEFAULT_INDENT = 4
40class ASTRewriter:
41 """AI: Rewriter that translates AST-level replace/remove/insert actions into byte-level edits."""
43 def __init__(
44 self,
45 node,
46 encoding: str = sys.getfilesystemencoding(),
47 correct_indent: bool = True,
48 ) -> None:
49 """AI: Accumulate and apply text rewrites (replace/insert/remove) to an AST node's source."""
50 self.__rewrites = _RewriteActions(node, encoding, correct_indent=correct_indent)
51 self.__filename = node.filename
53 def get_filename(self) -> str:
54 """AI: Return the filename of the AST node this rewriter operates on."""
55 return self.__filename
57 def replace(
58 self,
59 new_content: str,
60 target: Rewritable | Sequence[Rewritable] | PatternMatch | Sequence[PatternMatch],
61 include_whitespace: bool = True,
62 include_comments: bool = True,
63 ):
64 """AI: Queue a rewrite replacing target's source text with new_content."""
65 self.__rewrites.add(
66 _RewriteActionType.REPLACE,
67 target,
68 new_content,
69 include_whitespace,
70 include_comments,
71 )
73 def remove(
74 self,
75 target: Rewritable | Sequence[Rewritable] | PatternMatch | Sequence[PatternMatch],
76 include_whitespace: bool = True,
77 include_comments: bool = True,
78 ):
79 """AI: Queue a rewrite removing target's source text."""
80 self.__rewrites.add(_RewriteActionType.REMOVE, target, "", include_whitespace, include_comments)
82 def insert_before(
83 self,
84 new_content: str,
85 target: Rewritable | Sequence[Rewritable] | PatternMatch | Sequence[PatternMatch],
86 include_whitespace: bool = True,
87 include_comments: bool = True,
88 ):
89 """AI: Queue a rewrite inserting new_content immediately before target's source text."""
90 self.__rewrites.add(
91 _RewriteActionType.INSERT_BEFORE,
92 target,
93 new_content,
94 include_whitespace,
95 include_comments,
96 )
98 def insert_after(
99 self,
100 new_content: str,
101 target: Rewritable | Sequence[Rewritable] | PatternMatch | Sequence[PatternMatch],
102 include_whitespace: bool = True,
103 include_comments: bool = True,
104 ):
105 """AI: Queue a rewrite inserting new_content immediately after target's source text."""
106 self.__rewrites.add(
107 _RewriteActionType.INSERT_AFTER,
108 target,
109 new_content,
110 include_whitespace,
111 include_comments,
112 )
114 def apply_to_string(self) -> str:
115 """AI: Return the rewritten source as a string, applying all queued rewrites."""
116 return self.__rewrites.apply_to_string()
118 def apply(self) -> bytes:
119 """AI: Return the rewritten source as bytes, applying all queued rewrites (or the original content if none are queued)."""
120 if len(self.__rewrites.rewrites) == 0:
121 return self.__rewrites.content
122 return self.__rewrites.apply()
124 def has_changed(self) -> bool:
125 """AI: Return whether any rewrites have been queued."""
126 return len(self.__rewrites.rewrites) > 0
128 @staticmethod
129 def _get_comment_location(start_offset: int, stop_offset: int, content: bytes) -> tuple[int, int]:
130 return _RewriteActions.get_comment_location(start_offset, stop_offset, content)
133class _RewriteAction:
134 """Data container for a rewrite action to be applied later on to the AST."""
136 def __init__(
137 self,
138 action: _RewriteActionType,
139 target: Rewritable | Sequence[Rewritable] | PatternMatch | Sequence[PatternMatch],
140 replacement: str,
141 include_whitespace: bool,
142 include_comments: bool,
143 ) -> None:
144 self.action = action
145 self.target = target
146 self.replacement = replacement
147 self.nodes = self._get_nodes(target)
148 self.include_whitespace = include_whitespace
149 self.include_comments = include_comments
151 @staticmethod
152 def _get_nodes(
153 target: Rewritable | Sequence[Rewritable] | PatternMatch | Sequence[PatternMatch],
154 ) -> Sequence[Rewritable]:
155 if isinstance(target, Rewritable) or type(target).__name__ == "PythonASTNode":
156 return [target]
157 if isinstance(target, PatternMatch):
158 return target.nodes
159 assert isinstance(target, Sequence), "type of target violates its type requirements " + type(target).__name__
160 if len(target) > 0:
161 if isinstance(target[0], Rewritable): # TODO Why is part missing That is present on line 140, i.e.,
162 # or type(target).__name__ == "PythonASTNode"
163 return [n for n in target if isinstance(n, Rewritable)]
164 last = target[-1]
165 assert isinstance(last, PatternMatch), "type within Sequence violates its requirements " + type(last).__name__
166 return last.nodes
167 # TODO: is this correct? Can the other matches indeed be ignored?
168 return []
171class _RewriteActions:
172 """Data container for a list of rewrite actions to be applied later on to the AST."""
174 def __init__(
175 self,
176 node: Rewritable,
177 encoding: str,
178 correct_indent: bool,
179 rewrites: list[_RewriteAction] | None = None,
180 ) -> None:
181 self.rewrites: list[_RewriteAction] = rewrites or []
182 self.node = node
183 self.encoding = encoding
184 # self.content = self.node.root.binary_file_content()[self.node.offset : self.node.extended_end_offset]
185 self.content = node.text.encode(sys.getfilesystemencoding())
186 self.correct_indent = correct_indent
188 def add(
189 self,
190 action: _RewriteActionType,
191 target: Rewritable | Sequence[Rewritable] | PatternMatch | Sequence[PatternMatch],
192 replacement: str,
193 include_whitespace: bool,
194 include_comments: bool,
195 ):
196 rewrite = _RewriteAction(action, target, replacement, include_whitespace, include_comments)
197 self.add_rewrite(rewrite)
199 def add_rewrite(self, rewrite: _RewriteAction):
200 self.rewrites.append(rewrite)
202 def apply(self) -> bytes:
203 self.__check_for_conflicting_rewrites()
204 rewriter = Rewriter(self.content[:])
206 for rewrite in self.rewrites:
207 # skip nested rewrites as they are handled recursively by the parent rewrite
208 # except for if the rewrite node is the root node
209 if any(self.__is_ancestor_in_nodes(n) for n in rewrite.nodes if n != self.node):
210 continue
211 new_content, nodelist = self.__prepare_replacement_content(rewrite.replacement, rewrite.target)
212 if rewrite.action == _RewriteActionType.REPLACE:
213 self.__replace(
214 rewriter,
215 new_content,
216 nodelist,
217 rewrite.include_whitespace,
218 rewrite.include_comments,
219 )
220 elif rewrite.action == _RewriteActionType.INSERT_BEFORE:
221 self.__insert(
222 rewriter,
223 new_content,
224 True,
225 nodelist,
226 rewrite.include_whitespace,
227 rewrite.include_comments,
228 )
229 elif rewrite.action == _RewriteActionType.INSERT_AFTER:
230 self.__insert(
231 rewriter,
232 new_content,
233 False,
234 nodelist,
235 rewrite.include_whitespace,
236 rewrite.include_comments,
237 )
238 elif rewrite.action == _RewriteActionType.REMOVE:
239 self.__remove(
240 rewriter,
241 nodelist,
242 rewrite.include_whitespace,
243 rewrite.include_comments,
244 )
245 return rewriter.apply()
247 def apply_to_string(self) -> str:
248 return self.apply().decode(self.encoding)
250 def __check_for_conflicting_rewrites(self) -> None:
251 """Raise if two queued replace/remove rewrites target overlapping source ranges.
253 Insert rewrites are ignored: they only add text before or after a node without
254 overwriting it, so they can be combined with any other rewrite. A node nested inside
255 another rewritten node is not treated as overlapping.
257 TODO: this only turns silent corruption into a clear error - it doesn't merge
258 conflicting rewrites into a correct result. Recipes must still avoid queuing more than
259 one rewrite per node/range before a commit.
260 """
261 # INSERT_BEFORE/AFTER can work with each other and with REPLACE/REMOVE
262 replacing = [r for r in self.rewrites if r.action in (_RewriteActionType.REPLACE, _RewriteActionType.REMOVE)]
263 all_nodes = [(rewrite, node) for rewrite in replacing for node in rewrite.nodes if node != self.node]
264 for i, (rewrite_a, node_a) in enumerate(all_nodes):
265 for rewrite_b, node_b in all_nodes[i + 1 :]:
266 if rewrite_a is rewrite_b or self.__is_nested(node_a, node_b) or self.__is_nested(node_b, node_a):
267 continue
268 if not (node_a.end_offset < node_b.offset or node_a.offset > node_b.end_offset):
269 raise ValueError(
270 f"Conflicting rewrites queued for overlapping source ranges "
271 f"({node_a.offset}-{node_a.end_offset} and {node_b.offset}-{node_b.end_offset}) "
272 "in the same file - applying both would corrupt the output."
273 )
275 @staticmethod
276 def __is_nested(node: Rewritable, maybe_ancestor: Rewritable) -> bool:
277 """Return True if maybe_ancestor is an ancestor of node, walking .parent.
279 Not node.is_ancestor_of() - not every Rewritable implements it (e.g. PythonRstNode).
280 """
281 parent = node.parent
282 while parent:
283 if parent is maybe_ancestor:
284 return True
285 parent = parent.parent
286 return False
288 def __is_ancestor_in_nodes(self, node: Rewritable) -> bool:
289 """Check if the given node is a descendant of any nodes in the rewrite list.
291 Args:
292 node (Rewritable): The node to check.
294 Returns:
295 bool: True if the node is a descendant of any nodes in the rewrite list, False otherwise.
297 """
298 rewrite_nodes = list(flatten(rewrite.nodes for rewrite in self.rewrites))
300 # need to test
301 # 1
302 # | node |
303 # |rew|
304 # 2
305 # | rew |
306 # |node|
307 def no_conflict(node1, rew):
308 return not (node1.end_offset < rew.offset or node1.offset > rew.end_offset)
310 result = any(no_conflict(node, rew) for rew in rewrite_nodes)
312 return result and False # TODO: Why `and False`
314 def __replace(
315 self,
316 rewriter: Rewriter,
317 new_content: str,
318 nodes: Sequence[Rewritable],
319 include_whitespace: bool,
320 include_comments: bool,
321 ):
322 """Replace the content of the given node(s) with new content.
324 Args:
325 rewriter (Rewriter): The rewriter used to apply the content replacement.
326 new_content (str): The new content to insert in the specified range.
327 nodes (Sequence[Rewritable]): The nodes whose content is to be replaced.
328 include_whitespace (bool): Whether to include surrounding whitespace when determining the replacement range.
329 include_comments (bool): Whether to include surrounding comments when determining the replacement range.
331 """
332 if not nodes:
333 return
334 start_offset, end_offset = _RewriteActions.__correct_for_comments_and_whitespace(
335 self.node.offset,
336 self.content,
337 include_whitespace,
338 include_comments,
339 nodes,
340 )
341 # start_offset =nodes[0].get_start_offset()
342 # end_offset =nodes[-1].get_start_offset()+nodes[-1].get_length()+1
343 indent = self.derive_indent(start_offset)
344 if self.correct_indent:
345 if new_content.startswith("\n"):
346 # a blank first line would otherwise strand the original indent as trailing whitespace
347 start_offset -= indent
348 new_content = TextUtils.shift_right(new_content, indent, start_line=1)
349 self.__replace_bytes(rewriter, start_offset, end_offset, new_content)
351 def __remove(
352 self,
353 rewriter: Rewriter,
354 nodes: Sequence[Rewritable],
355 include_whitespace: bool = False,
356 include_comments: bool = False,
357 ):
358 """Remove a list of AST nodes from the content, optionally including surrounding whitespace and comments.
360 Args:
361 rewriter (Rewriter): The rewriter used to apply the removal.
362 nodes (Sequence[Rewritable]): The list of AST nodes to remove.
363 include_whitespace (bool, optional): Whether to include surrounding whitespace in the removal. Defaults to False.
364 include_comments (bool, optional): Whether to include surrounding comments in the removal. Defaults to False.
366 Returns:
367 None
369 """
370 if not nodes:
371 return
373 start_offset, end_offset = _RewriteActions.__correct_for_comments_and_whitespace(
374 self.node.offset,
375 self.content,
376 include_whitespace,
377 include_comments,
378 nodes,
379 )
380 indent = self.derive_indent(start_offset)
381 # remove the indent in front of it
382 start_offset -= indent
383 # remove the line if it is empty
384 if (
385 start_offset > 0
386 and self.content[start_offset - 1] == ord("\n")
387 and (end_offset >= len(self.content) or self.content[end_offset] == ord("\n"))
388 ):
389 start_offset -= 1
390 self.__replace_bytes(rewriter, start_offset, end_offset, "")
392 def derive_indent(self, start_offset: int) -> int:
393 indent = 0 # len(nodes[0].indent)
394 if start_offset > 0:
395 while len(self.content) > (start_offset - indent - 1) and self.content[start_offset - indent - 1] == 32:
396 indent += 1
397 return indent
399 def __insert(
400 self,
401 rewriter: Rewriter,
402 new_content: str,
403 before: bool,
404 nodes: Sequence[Rewritable],
405 include_whitespace: bool,
406 include_comments: bool,
407 ):
408 if not nodes:
409 return
410 content = self.content
411 indent = TextUtils.get_spaces_before(content, nodes[0].offset)
412 spaces = " " * indent
413 # if flattened_nodes[-1] has a new line after white space then we need to add a new line:
414 ext_start_offset, ext_end_offset = _RewriteActions.__correct_for_comments_and_whitespace(
415 self.node.offset,
416 self.content,
417 include_whitespace,
418 include_comments,
419 nodes,
420 )
421 white_space = (
422 ""
423 if not include_whitespace
424 else "\n" + spaces
425 if ext_end_offset < len(content) and content[ext_end_offset] in b"\n"
426 else spaces
427 )
428 # indent the new content except the first line
429 new_content = TextUtils.shift_right(new_content, indent, start_line=1)
431 if before:
432 # restore the node's original indent, consumed as new_content's first-line indent
433 if not white_space and new_content.endswith("\n"):
434 new_content += spaces
435 self.__replace_bytes(rewriter, ext_start_offset, ext_start_offset, new_content + white_space)
436 else:
437 self.__replace_bytes(rewriter, ext_end_offset, ext_end_offset, white_space + new_content)
439 def __replace_bytes(self, rewriter: Rewriter, start: int, end: int, new_content: str) -> None:
440 """Replace the content in the specified range with new content.
442 Args:
443 rewriter (Rewriter): The rewriter used to apply the byte replacement.
444 start (int): The starting index of the range to be replaced.
445 end (int): The ending index of the range to be replaced.
446 new_content (str): The new content to insert in the specified range.
448 """
449 rewriter.replace(start, end, new_content.encode(self.encoding))
451 def __compose_replacement(self, replacement: str, matches: Sequence[PatternMatch]) -> str:
452 all_placeholders = {p: n for m in matches for p, n in m.expansions.items()}
453 for placeholder, nodes in all_placeholders.items():
454 quoted_placeholder = re.escape(placeholder)
455 raw_signature = self.__get_texts(nodes)
456 # replacement = replacement.replace(placeholder, raw_signature)
457 while placeholder in replacement:
458 pattern = re.compile(r"( *)" + quoted_placeholder)
459 matcher = pattern.search(replacement)
461 if matcher:
462 spaces = matcher[1]
463 place_holder_length = len(placeholder)
464 index = replacement.index(placeholder)
465 # TODO a regex may be provided between backticks and the groups are used. This needs a better design
466 # A preferable solution is to pass a transformer function to the compose_replacement
467 if index + place_holder_length < len(replacement) and replacement[index + place_holder_length] == "`":
468 # ` ` means get regex
469 end_index = replacement.index("`", index + place_holder_length + 1)
470 if not end_index:
471 raise ValueError("No closing ` found")
472 regex = replacement[index + place_holder_length + 1 : end_index]
473 regex_match = re.match(regex, raw_signature)
474 if regex_match:
475 raw_signature = "".join(regex_match.groups())
476 place_holder_length = end_index - index + 1
477 indent_replacement = raw_signature.replace("\n", "\n" + spaces)
478 if (
479 placeholder.startswith("$$")
480 and index + place_holder_length < len(replacement)
481 and replacement[index + place_holder_length] == ";"
482 ):
483 place_holder_length += 1
484 # replace the placeholder with the indent replacement
485 replacement = replacement[:index] + indent_replacement + replacement[index + place_holder_length :]
486 else:
487 print("Match doesn't match unexpectedly")
488 return replacement
490 def __get_texts(self, nodes: Sequence[Rewritable]) -> str:
491 if len(nodes) == 1:
492 return self.__get_text(nodes[0])
493 # Use a ASTRewriter to only rewrite exactly that what needs to be rewritten
494 rewriter = ASTRewriter(nodes[0], self.encoding, correct_indent=False)
495 for node in nodes:
496 rs = self.__get_text(node)
497 org_rs = node.text
498 if rs != org_rs:
499 rewriter.replace(rs, node)
500 result = rewriter.apply_to_string()
501 indent = self.derive_indent(nodes[0].offset)
502 return TextUtils.shift_left(result, indent, start_line=1)
504 def __get_text(self, node: Rewritable) -> str:
505 if self._should_skip(node):
506 return ""
508 if node == self.node:
509 return node.text
510 # the descendants may need to be rewritten as well
511 # rewrites = [rewrite for rewrite in self.rewrites if any(node.is_ancestor_of(rewrite_node)
512 # for rewrite_node in rewrite.nodes)]
513 rewrites = [rewrite for rewrite in self.rewrites if any(node.is_ancestor_of(rewrite_node) for rewrite_node in rewrite.nodes)]
514 if rewrites:
515 rewriter = _RewriteActions(node, self.encoding, self.correct_indent, rewrites)
516 return rewriter.apply_to_string()
517 return node.text
519 def __prepare_replacement_content(
520 self,
521 new_content: str,
522 target: PatternMatch | Rewritable | Sequence[Rewritable],
523 ) -> tuple[str, Sequence[Rewritable]]:
524 if isinstance(target, PatternMatch):
525 new_content = self.__compose_replacement(new_content, [target])
526 node_list = target.nodes
527 else:
528 node_list = (
529 [target] if (isinstance(target, Rewritable) or type(target).__name__ == "PythonASTNode") else target
530 ) # TODO How to make a Sequence[Rewritable] as type hints also show list[Rewritable]?
531 return new_content, node_list
533 def _should_skip(self, node: Rewritable):
534 """If the node is not the first node of a pattern match it should be skipped."""
535 return any(node in rewrite.nodes[1:] for rewrite in self.rewrites if isinstance(rewrite.target, PatternMatch))
537 @staticmethod
538 def _get_parent_statement(node: Rewritable):
539 parent = node
540 while parent and not parent.is_statement:
541 parent = parent.parent
542 return parent
544 @staticmethod
545 def __correct_for_comments_and_whitespace(
546 offset: int,
547 content: bytes,
548 include_whitespace: bool,
549 include_comments: bool,
550 nodes: Sequence[Rewritable],
551 ):
552 start_offset = nodes[0].offset - offset
553 end_offset = nodes[-1].extended_end_offset - offset
554 if include_comments:
555 preceding_node = nodes[0].preceding_sibling
556 parent = nodes[0].parent
557 start_comment_location = 0
558 if preceding_node:
559 # start after the comment of the preceding node
560 start_comment_location = preceding_node.extended_end_offset - offset
561 preceding_end_offset = _RewriteActions.__get_comment_after_location(start_comment_location, start_offset, content)
562 if preceding_end_offset != (-1, -1):
563 start_comment_location = preceding_end_offset[1]
564 elif parent:
565 start_comment_location = parent.offset - offset
566 # get the comment belonging to the preceding node
567 extended_location = _RewriteActions.get_comment_location(start_comment_location, start_offset, content)
568 if extended_location != (-1, -1):
569 start_offset = extended_location[0]
570 next_sibling = nodes[-1].next_sibling
571 end_comment_location = next_sibling.offset - offset if next_sibling else parent.end_offset - offset if parent else len(content)
572 location_after_comment = _RewriteActions.__get_comment_after_location(end_offset, end_comment_location, content)
573 if location_after_comment != (-1, -1):
574 end_offset = location_after_comment[1]
575 if include_whitespace:
576 end_offset = _RewriteActions.__extend_with_whitespace(end_offset, content)
577 return start_offset, end_offset
579 def cor_offset(self, offset: int):
580 return offset - self.node.offset
582 @staticmethod
583 def get_comment_location(start_offset: int, stop_offset: int, content: bytes) -> tuple[int, int]:
584 """Get the location of the comment before the location, but after the stop_location.
586 A comment is a line that starts with // or a block that starts with /* and ends with */
587 or a line that starts with #.
588 """
589 # TODO: the // and # branches below only find the single closest comment line
590 # search last occurrence of //, /*, # in a byte array
591 comment_start = content.rfind(b"//", start_offset, stop_offset)
592 if comment_start != -1:
593 comment_end = _RewriteActions.__get_end_of_line(content, comment_start)
594 return comment_start, comment_end
595 comment_start = content.rfind(b"/*", start_offset, stop_offset)
596 if comment_start != -1:
597 comment_end = content.find(b"*/", comment_start, stop_offset)
598 if comment_end != -1:
599 comment_end += len("*/")
600 return comment_start, comment_end
601 comment_start = content.rfind(b"#", start_offset, stop_offset)
602 if comment_start != -1:
603 comment_end = _RewriteActions.__get_end_of_line(content, comment_start)
604 return comment_start, comment_end
605 return -1, -1
607 @staticmethod
608 def __extend_with_whitespace(start_offset: int, content: bytes) -> int:
609 end_location = _RewriteActions.__get_end_of_line(content, start_offset)
610 text = content[start_offset:end_location]
611 for byt in text:
612 if byt not in b" \t":
613 return start_offset
614 return end_location
616 @staticmethod
617 def __get_comment_after_location(start_offset: int, end_offset: int, content: bytes) -> tuple[int, int]:
618 """Get the location of the comment before the location, but after the stop_location.
620 A comment is a line that starts with // or a block that starts with /* and ends with */
621 or a line that starts with #.
622 """
623 # TODO: same single-line limitation as get_comment_location - a multi-line trailing
624 # comment block is only captured up to its first line here.
625 line_end_offset = _RewriteActions.__get_end_of_line(content, start_offset)
626 if line_end_offset == -1:
627 line_end_offset = len(content)
628 comment_start = content.find(b"//", start_offset, line_end_offset)
629 if comment_start == -1:
630 comment_start = content.rfind(b"#", start_offset, line_end_offset)
631 if comment_start != -1:
632 return comment_start, line_end_offset
633 comment_start = content.rfind(b"/*", start_offset, line_end_offset)
634 if comment_start != -1:
635 # a block comment must start on the same line but doesn't have to finish on the same line
636 comment_end = content.find(b"*/", comment_start, end_offset)
637 if comment_end != -1:
638 comment_end += len("*/")
639 return comment_start, comment_end
640 return -1, -1
642 @staticmethod
643 def __get_end_of_line(content: bytes, start: int):
644 location = content.find(b"\n", start)
645 if location == -1:
646 return len(content)
647 return location
649 @staticmethod
650 def __get_depth(node: Rewritable) -> int:
651 depth = 0
652 parent = node.parent
653 while parent:
654 if parent.semantic_kind is SemanticKind.STATEMENT or parent.parser_kind in {"CompoundStmt", "COMPOUND_STMT"}:
655 depth += 1
656 parent = parent.parent
657 return depth