Fix mistyped command heads in every pipe segment

This commit is contained in:
Alexander
2026-09-13 20:49:04 +02:00
parent 624938c284
commit a9e45d0186
2 changed files with 86 additions and 21 deletions
+59
View File
@@ -1,4 +1,6 @@
import pytest
from thefuck import shell_ast
from thefuck.learned import LearnedCorrections
@@ -160,6 +162,63 @@ class TestGuessFromPath(object):
assert learned.guess_from_path('') is None
class TestGuessFromPathSegments(object):
pytestmark = pytest.mark.skipif(
not shell_ast.AST_AVAILABLE, reason='bashlex required')
@pytest.fixture
def path_bins(self, monkeypatch):
def setup(executables, existing=()):
monkeypatch.setattr('thefuck.learned.get_all_executables',
lambda: list(executables))
monkeypatch.setattr('thefuck.learned.which',
lambda token: token in existing)
return setup
def test_fixes_head_of_every_pipe_segment(self, learned, path_bins):
path_bins(executables=['git', 'grep', 'sed'])
assert (learned.guess_from_path('gi psuh | gre -i foo')
== 'git psuh | grep -i foo')
def test_fixes_only_segment_with_unknown_head(self, learned, path_bins):
path_bins(executables=['git', 'grep', 'sed'], existing=['git'])
assert (learned.guess_from_path('git psuh | gre -i foo')
== 'git psuh | grep -i foo')
def test_returns_none_when_all_heads_executable(self, learned, path_bins):
path_bins(executables=['git', 'grep'], existing=['git', 'grep'])
assert learned.guess_from_path('git psuh | grep -i foo') is None
def test_skips_ambiguous_segment_fixes_others(self, learned, path_bins):
path_bins(executables=['clear', 'clean', 'grep'])
assert (learned.guess_from_path('clea psuh | gre -i foo')
== 'clea psuh | grep -i foo')
def test_guesses_after_sudo_in_pipe(self, learned, path_bins):
path_bins(executables=['clear', 'grep'])
assert (learned.guess_from_path('sudo cler | gre -i foo')
== 'sudo clear | grep -i foo')
def test_guesses_after_sudo_flat(self, learned, path_bins):
path_bins(executables=['clear', 'grep', 'sed'])
assert learned.guess_from_path('sudo cler') == 'sudo clear'
def test_replaces_quoted_head_whole(self, learned, path_bins):
path_bins(executables=['clear', 'grep', 'sed'])
assert learned.guess_from_path('"cler" psuh') == 'clear psuh'
def test_preserves_spacing_between_words(self, learned, path_bins):
path_bins(executables=['clear', 'grep', 'sed'])
assert learned.guess_from_path('cler psuh') == 'clear psuh'
def test_unparseable_script_falls_back_to_flat_view(
self, learned, path_bins):
path_bins(executables=['clear', 'grep', 'sed'])
assert (
learned.guess_from_path('cler psuh; case $x in y) ;; esac')
== 'clear psuh; case $x in y) ;; esac')
class TestClear(object):
def test_removes_all_entries(self, learned):
learned.record("git psuh", "git push")
+27 -21
View File
@@ -4,7 +4,7 @@ import shelve
import time
from difflib import get_close_matches
from . import logs
from . import logs, shell_ast
from .utils import get_all_executables, which
try:
@@ -144,29 +144,35 @@ class LearnedCorrections(object):
def guess_from_path(self, script):
"""Guesses what the user meant by fuzzy-matching a mistyped
token against executables from $PATH and shell aliases.
command token against executables from $PATH and shell aliases,
for the head of every pipeline segment, so a typo after a pipe
is fixed as reliably as a leading one.
Returns the corrected script only when exactly one unambiguous
close match exists, so asking the user stays the last resort.
Returns the corrected script only when at least one segment
has exactly one unambiguous close match, so asking the user
stays the last resort.
"""
parts = script.split()
if not parts:
replacements = []
for segment in shell_ast.parse(script):
words = segment.words
index = 1 if words[0][0] == 'sudo' and len(words) > 1 else 0
token, start, end = words[index]
# Raw spans keep their quotes; the gates must see the typed
# word while the whole quoted span is replaced below.
if (len(token) > 1 and token[0] == token[-1]
and token[0] in ('"', "'")):
token = token[1:-1]
if not token or '/' in token or '.' in token or which(token):
continue
candidates = [cmd for cmd in get_close_matches(
token, get_all_executables(), n=5, cutoff=GUESS_CUTOFF)
if cmd.startswith(token[0])]
if len(candidates) != 1:
continue
replacements.append((start, end, candidates[0]))
if not replacements:
return None
index = 1 if len(parts) > 1 and parts[0] == 'sudo' else 0
token = parts[index]
if '/' in token or '.' in token or which(token):
return None
candidates = [cmd for cmd in get_close_matches(
token, get_all_executables(), n=5, cutoff=GUESS_CUTOFF)
if cmd.startswith(token[0])]
if len(candidates) != 1:
return None
corrected_parts = list(parts)
corrected_parts[index] = candidates[0]
return ' '.join(corrected_parts)
return shell_ast.splice(script, replacements)
def clear(self):
db = self.db