From a5c675cdeac7d92c44cf1ee433312a32faa4caf3 Mon Sep 17 00:00:00 2001 From: Alexander Date: Sun, 13 Sep 2026 20:39:08 +0200 Subject: [PATCH] Add shell AST core with pos-splice round-trip parse() turns a script into pipeline command segments via bashlex, keeping byte-true char offsets for every word (quotes included as typed, redirects and env assignments excluded from words), and splice() rebuilds a corrected script from (start, end, text) replacements applied right-to-left. When bashlex is missing or rejects the script, a flat split(' ') segment keeps the old view. --- tests/test_ast_core.py | 198 +++++++++++++++++++++++++++++++++++++++++ thefuck/shell_ast.py | 105 ++++++++++++++++++++++ 2 files changed, 303 insertions(+) create mode 100644 tests/test_ast_core.py diff --git a/tests/test_ast_core.py b/tests/test_ast_core.py new file mode 100644 index 0000000..0470fd6 --- /dev/null +++ b/tests/test_ast_core.py @@ -0,0 +1,198 @@ +import pytest + +from thefuck import shell_ast +from thefuck.shell_ast import parse, splice + + +needs_bashlex = pytest.mark.skipif( + not shell_ast.AST_AVAILABLE, reason='bashlex is not available') + + +CORPUS = [ + 'gti psuh | greo -i foo', + 'ls', + 'one | two | three', + 'a && b', + 'a; b', + "echo 'hello world' foo", + 'echo "double quoted arg"', + 'cmd > file 2>&1', + 'FOO=bar cmd arg', + 'sudo apt-get install vim', + 'trailing spaces ', + 'ls | where size > 1mb', + 'echo naïve', + 'cmd arg', + 'cd /tmp && ls -la', + 'grep -lir . test | sort | uniq', + 'python -c "print(1)"', + 'echo a"b c"d', + 'sudo docker ps -a | grep exited', + 'echo one > file', +] + +# Each entry falls back because bashlex 0.18 raises for it: +# `case` and $((...)) hit NotImplementedError, empty and whitespace-only +# scripts hit its empty-string AttributeError, the lone pipe is a +# ParsingError, and the unclosed quote raises MatchedPairError, +# which subclasses bashlex.errors.ParsingError. +FALLBACK_SCRIPTS = [ + 'case $x in a) foo;; esac', + 'echo $((1+1))', + '', + ' ', + '|', + "'abc", +] + +EXPECTED_HEADS = [ + ('gti psuh | greo -i foo', ['gti', 'greo']), + ('ls', ['ls']), + ('one | two | three', ['one', 'two', 'three']), + ('a && b', ['a', 'b']), + ('a; b', ['a', 'b']), + ('trailing spaces ', ['trailing']), + ('ls | where size > 1mb', ['ls', 'where']), + ('cd /tmp && ls -la', ['cd', 'ls']), + ('grep -lir . test | sort | uniq', ['grep', 'sort', 'uniq']), + ('sudo docker ps -a | grep exited', ['sudo', 'grep']), + ('echo one > file', ['echo']), +] + + +@pytest.mark.parametrize('script', CORPUS) +def test_splice_with_no_replacements_is_identity(script): + assert splice(script, []) == script + + +@needs_bashlex +@pytest.mark.parametrize('script', CORPUS) +def test_word_offsets_round_trip_byte_for_byte(script): + segments = parse(script) + for segment in segments: + for token, start, end in segment.words: + assert script[start:end] == token + assert splice(script, [(start, end, token)]) == script + everything = [(start, end, token) + for segment in segments + for token, start, end in segment.words] + assert splice(script, everything) == script + + +@needs_bashlex +def test_replaces_misspelled_heads_across_pipeline(): + script = 'gti psuh | greo -i foo' + replacements = [] + for segment in parse(script): + if segment.head == 'gti': + replacements.append( + (segment.head_start, segment.head_end, 'git')) + elif segment.head == 'greo': + replacements.append( + (segment.head_start, segment.head_end, 'grep')) + assert splice(script, replacements) == 'git psuh | grep -i foo' + + +@needs_bashlex +@pytest.mark.parametrize('script,heads', EXPECTED_HEADS) +def test_segment_heads(script, heads): + assert [segment.head for segment in parse(script)] == heads + + +@needs_bashlex +def test_pipeline_segments_carry_word_and_span_offsets(): + first, second = parse('gti psuh | greo -i foo') + assert first.words == [('gti', 0, 3), ('psuh', 4, 8)] + assert (first.head_start, first.head_end) == (0, 3) + assert (first.start, first.end) == (0, 8) + assert second.words == [('greo', 11, 15), ('-i', 16, 18), + ('foo', 19, 22)] + assert (second.start, second.end) == (11, 22) + + +@needs_bashlex +def test_env_prefix_head_is_the_real_command(): + # bashlex 0.18 classifies `FOO=bar` as an assignment node, not a + # word, so the first collected word is the command itself and the + # head resolves to `cmd`. + segment, = parse('FOO=bar cmd arg') + assert segment.head == 'cmd' + assert segment.words == [('cmd', 8, 11), ('arg', 12, 15)] + + +@needs_bashlex +def test_redirect_words_are_excluded_from_segments(): + segment, = parse('cmd > file 2>&1') + assert segment.words == [('cmd', 0, 3)] + + +@needs_bashlex +def test_nushell_comparison_stays_byte_true(): + # `> 1mb` misparses as a redirect node; it must stay out of the + # words while the remaining offsets keep addressing the original. + first, second = parse('ls | where size > 1mb') + assert first.words == [('ls', 0, 2)] + assert second.words == [('where', 5, 10), ('size', 11, 15)] + + +@needs_bashlex +def test_quoted_word_spans_include_the_quotes(): + segment, = parse("echo 'hello world' foo") + assert segment.words == [('echo', 0, 4), ("'hello world'", 5, 18), + ('foo', 19, 22)] + assert splice("echo 'hello world' foo", + [(5, 18, "'hi'")]) == "echo 'hi' foo" + + +@needs_bashlex +def test_mixed_quoted_word_splices_as_one_span(): + segment, = parse('echo a"b c"d') + assert segment.words == [('echo', 0, 4), ('a"b c"d', 5, 12)] + assert splice('echo a"b c"d', [(5, 12, "'x y'")]) == "echo 'x y'" + + +@needs_bashlex +@pytest.mark.parametrize('script', FALLBACK_SCRIPTS) +def test_unparseable_scripts_fall_back_flat(script): + segments = parse(script) + assert len(segments) == 1 + segment = segments[0] + assert [word[0] for word in segment.words] == script.split(' ') + for token, start, end in segment.words: + assert script[start:end] == token + assert (segment.start, segment.end) == (0, len(script)) + + +def test_flat_fallback_when_ast_disabled(monkeypatch): + monkeypatch.setattr(shell_ast, 'AST_AVAILABLE', False) + script = 'gti psuh | greo -i foo' + segment, = parse(script) + assert segment.head == 'gti' + assert [word[0] for word in segment.words] == script.split(' ') + for token, start, end in segment.words: + assert script[start:end] == token + assert (segment.start, segment.end) == (0, len(script)) + + +def test_flat_fallback_offsets_survive_consecutive_spaces(monkeypatch): + monkeypatch.setattr(shell_ast, 'AST_AVAILABLE', False) + script = 'gti psuh' + segment, = parse(script) + # 'gti psuh'.split(' ') keeps an empty token between the spaces + assert segment.words == [('gti', 0, 3), ('', 4, 4), ('psuh', 5, 9)] + assert splice(script, [(0, 3, 'git')]) == 'git psuh' + + +def test_splice_sorts_replacements_and_applies_right_to_left(): + assert splice('abcdef', [(4, 6, 'Y'), (0, 3, 'X')]) == 'XdY' + + +def test_splice_allows_touching_replacements(): + assert splice('abcdef', [(0, 3, 'X'), (3, 6, 'YZ')]) == 'XYZ' + + +def test_splice_rejects_overlapping_replacements(): + with pytest.raises(ValueError): + splice('abcdef', [(0, 3, 'X'), (2, 5, 'Y')]) + with pytest.raises(ValueError): + splice('abcdef', [(0, 3, 'X'), (0, 3, 'Y')]) diff --git a/thefuck/shell_ast.py b/thefuck/shell_ast.py index c4f8907..04faa5c 100644 --- a/thefuck/shell_ast.py +++ b/thefuck/shell_ast.py @@ -12,3 +12,108 @@ except ImportError: AST_AVAILABLE = bashlex is not None and six.PY3 + +# bashlex 0.18 raises bare NotImplementedError for constructs it does +# not support (like `case` and arithmetic expansion) and AttributeError +# for empty or whitespace-only scripts; MatchedPairError, raised for +# unclosed quotes, subclasses ParsingError. +_FALLBACK_ERRORS = (AttributeError, NotImplementedError, TypeError) +if bashlex is not None: + _FALLBACK_ERRORS += (bashlex.errors.ParsingError,) + + +class Segment(object): + """One command of a pipeline with its words' char offsets. + + Every offset addresses the original script string, and every + token is the raw span text (`script[start:end]`, quotes included + as typed), so replacing a span with its token is always an + identity and a corrected script is rebuilt by splicing + replacements in at those offsets. + """ + + def __init__(self, words, start, end): + self.words = words + head, head_start, head_end = words[0] + self.head = head + self.head_start = head_start + self.head_end = head_end + self.start = start + self.end = end + + +def parse(script): + """Splits a script into pipeline command segments. + + Falls back to a single flat segment when bashlex is missing or + cannot parse the script, so callers always get a usable view. + """ + if not AST_AVAILABLE: + return [_flat_segment(script)] + try: + trees = bashlex.parse(script) + except _FALLBACK_ERRORS: + return [_flat_segment(script)] + collector = _SegmentCollector(script) + for tree in trees: + collector.visit(tree) + if not collector.segments: + return [_flat_segment(script)] + return collector.segments + + +def splice(script, replacements): + """Rebuilds a script from (start, end, new_text) replacements. + + Replacements are applied right-to-left so earlier offsets stay + valid; overlapping spans raise ValueError because callers must + never overlap. + """ + ordered = sorted(replacements, key=lambda replacement: replacement[0]) + for index in range(1, len(ordered)): + if ordered[index][0] < ordered[index - 1][1]: + raise ValueError('overlapping replacements: {} and {}'.format( + ordered[index - 1], ordered[index])) + for start, end, text in reversed(ordered): + script = script[:start] + text + script[end:] + return script + + +def _flat_segment(script): + """Builds the fallback segment from space-split words. + + Mirrors the flat view of `learned.get_correction`: offsets follow + `script.split(' ')` exactly, so consecutive spaces yield empty + tokens with zero-width spans at their true positions. + """ + words = [] + start = 0 + for token in script.split(' '): + end = start + len(token) + words.append((token, start, end)) + start = end + 1 + return Segment(words, 0, len(script)) + + +if AST_AVAILABLE: + class _SegmentCollector(bashlex.ast.nodevisitor): + """Collects command nodes as segments. + + Redirect and assignment parts are skipped, so their words + never enter a segment; nested commands (inside substitutions + or compounds of an already collected command) are skipped the + same way by not descending into a command's parts. + """ + + def __init__(self, script): + self.script = script + self.segments = [] + + def visitcommand(self, node, parts): + words = [(self.script[part.pos[0]:part.pos[1]], + part.pos[0], part.pos[1]) + for part in node.parts if part.kind == 'word'] + if words: + self.segments.append( + Segment(words, node.pos[0], node.pos[1])) + return False