diff --git a/tests/entrypoints/test_fix_command_learned.py b/tests/entrypoints/test_fix_command_learned.py index a5f22e1..26b1a02 100644 --- a/tests/entrypoints/test_fix_command_learned.py +++ b/tests/entrypoints/test_fix_command_learned.py @@ -1,23 +1,31 @@ import pytest -from mock import Mock, patch, call +from mock import Mock, patch from thefuck.entrypoints.fix_command import fix_command from thefuck.types import CorrectedCommand @pytest.fixture def mock_learned(monkeypatch): - state = {"correction": None, "recordings": []} + state = {"correction": None, "guess": None, "recordings": []} def fake_get_correction(script): return state["correction"] + def fake_guess(script): + return state["guess"] + def fake_record(original, corrected): state["recordings"].append((original, corrected)) monkeypatch.setattr( "thefuck.entrypoints.fix_command.get_correction", fake_get_correction ) - monkeypatch.setattr("thefuck.entrypoints.fix_command.record", fake_record) + monkeypatch.setattr( + "thefuck.entrypoints.fix_command.guess_from_path", fake_guess + ) + monkeypatch.setattr( + "thefuck.entrypoints.fix_command.record", fake_record + ) return state @@ -75,6 +83,41 @@ class TestLearnedAutoApply(object): assert shown_cmd.script == "git push origin main" +class TestGuessAutoApply(object): + def test_guess_records_and_auto_applies( + self, mock_learned, known_args, settings, monkeypatch + ): + mock_learned["guess"] = "git push origin main" + get_corrected = Mock() + monkeypatch.setattr( + "thefuck.entrypoints.fix_command.get_corrected_commands", get_corrected + ) + + with patch("thefuck.types.CorrectedCommand.run") as mock_run, patch( + "thefuck.logs.show_corrected_command" + ): + fix_command(known_args) + mock_run.assert_called_once() + get_corrected.assert_not_called() + assert mock_learned["recordings"] == [ + ("git psuh origin main", "git push origin main") + ] + + def test_exact_learned_wins_over_guess( + self, mock_learned, known_args, settings, monkeypatch + ): + mock_learned["correction"] = "git push --force origin main" + mock_learned["guess"] = "git push origin main" + + with patch("thefuck.types.CorrectedCommand.run"), patch( + "thefuck.logs.show_corrected_command" + ) as mock_show: + fix_command(known_args) + shown_cmd = mock_show.call_args[0][0] + assert shown_cmd.script == "git push --force origin main" + assert mock_learned["recordings"] == [] + + class TestRecordOnSelection(object): def test_records_user_selection( self, mock_learned, known_args, settings, monkeypatch diff --git a/tests/test_learned.py b/tests/test_learned.py index c9bab33..0486b1d 100644 --- a/tests/test_learned.py +++ b/tests/test_learned.py @@ -108,6 +108,58 @@ class TestGetCorrection(object): assert learned.get_correction("sl") == "ls" +class TestGuessFromPath(object): + @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_guesses_unique_close_match(self, learned, path_bins): + path_bins(executables=['clear', 'grep', 'sed']) + assert learned.guess_from_path('cler') == 'clear' + + def test_keeps_arguments(self, learned, path_bins): + path_bins(executables=['python', 'pydoc', 'grep']) + assert (learned.guess_from_path('pyhton script.py') + == 'python script.py') + + def test_returns_none_when_ambiguous(self, learned, path_bins): + path_bins(executables=['clear', 'clean']) + assert learned.guess_from_path('clea') is None + + def test_returns_none_when_token_is_executable(self, learned, path_bins): + path_bins(executables=['clear'], existing=['clear']) + assert learned.guess_from_path('clear') is None + + def test_returns_none_for_path_like_token(self, learned, path_bins): + path_bins(executables=['git', 'grep', 'sed']) + assert learned.guess_from_path('./gti push') is None + + def test_returns_none_for_token_with_extension(self, learned, path_bins): + path_bins(executables=['git', 'grep', 'sed']) + assert learned.guess_from_path('giti.py x') is None + + def test_returns_none_when_first_char_differs(self, learned, path_bins): + path_bins(executables=['top']) + assert learned.guess_from_path('htop') is None + + def test_returns_none_below_cutoff(self, learned, path_bins): + path_bins(executables=['grep', 'sed', 'awk']) + assert learned.guess_from_path('total') is None + + def test_guesses_after_sudo(self, learned, path_bins): + path_bins(executables=['clear', 'grep', 'sed']) + assert learned.guess_from_path('sudo cler') == 'sudo clear' + + def test_returns_none_for_empty_script(self, learned, path_bins): + path_bins(executables=['git']) + assert learned.guess_from_path('') is None + + class TestClear(object): def test_removes_all_entries(self, learned): learned.record("git psuh", "git push") diff --git a/thefuck/entrypoints/fix_command.py b/thefuck/entrypoints/fix_command.py index 9722aee..22bdd15 100644 --- a/thefuck/entrypoints/fix_command.py +++ b/thefuck/entrypoints/fix_command.py @@ -6,7 +6,7 @@ from .. import logs, types, const from ..conf import settings from ..corrector import get_corrected_commands from ..exceptions import EmptyCommand -from ..learned import get_correction, record +from ..learned import get_correction, guess_from_path, record from ..ui import select_command from ..utils import get_alias, get_all_executables @@ -41,6 +41,12 @@ def fix_command(known_args): return learned_script = get_correction(command.script) + if not learned_script: + learned_script = guess_from_path(command.script) + if learned_script: + logs.debug("Guessed correction from $PATH: {}".format( + learned_script)) + record(command.script, learned_script) if learned_script: learned_cmd = types.CorrectedCommand( script=learned_script, side_effect=None, priority=0 diff --git a/thefuck/learned.py b/thefuck/learned.py index ec74b13..3722f9d 100644 --- a/thefuck/learned.py +++ b/thefuck/learned.py @@ -2,8 +2,10 @@ import atexit import os import shelve import time +from difflib import get_close_matches from . import logs +from .utils import get_all_executables, which try: import dbm @@ -18,6 +20,9 @@ except ImportError: _shelve_open_error = () +GUESS_CUTOFF = 0.8 + + class LearnedCorrections(object): def __init__(self): self._db = None @@ -137,6 +142,32 @@ class LearnedCorrections(object): return None + def guess_from_path(self, script): + """Guesses what the user meant by fuzzy-matching a mistyped + token against executables from $PATH and shell aliases. + + Returns the corrected script only when exactly one unambiguous + close match exists, so asking the user stays the last resort. + """ + parts = script.split() + if not parts: + 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) + def clear(self): db = self.db for key in list(db.keys()): @@ -154,4 +185,5 @@ _learned = LearnedCorrections() record = _learned.record get_correction = _learned.get_correction +guess_from_path = _learned.guess_from_path clear = _learned.clear