diff --git a/tests/test_learned.py b/tests/test_learned.py index 0486b1d..7fa4f01 100644 --- a/tests/test_learned.py +++ b/tests/test_learned.py @@ -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") diff --git a/thefuck/learned.py b/thefuck/learned.py index 3722f9d..b22809c 100644 --- a/thefuck/learned.py +++ b/thefuck/learned.py @@ -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