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.
This commit is contained in:
Alexander
2026-09-13 20:39:08 +02:00
parent f7d3b10288
commit a5c675cdea
2 changed files with 303 additions and 0 deletions
+198
View File
@@ -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')])
+105
View File
@@ -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