From 1d3e4b7e8ade38a1a0c8cca4e51ee73d3daaa7f4 Mon Sep 17 00:00:00 2001 From: Alexander Date: Sun, 13 Sep 2026 20:01:22 +0200 Subject: [PATCH] Guess mistyped commands from $PATH before asking the user When no learned correction matches, fuzzy-match the mistyped token against executables from $PATH and shell aliases. A guess is auto-run and recorded into the learned-corrections db only when it is unambiguous: the token is not an existing executable, contains no path separators or extensions, shares its first character with the candidate, has similarity of at least 0.8, and exactly one candidate survives. Asking the user stays the last resort. Ultraworked with [Sisyphus](https://github.com/code-yeongyu/oh-my-openagent) Co-authored-by: Sisyphus --- tests/entrypoints/test_fix_command_learned.py | 49 +++++++++++++++-- tests/test_learned.py | 52 +++++++++++++++++++ thefuck/entrypoints/fix_command.py | 8 ++- thefuck/learned.py | 32 ++++++++++++ 4 files changed, 137 insertions(+), 4 deletions(-) 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