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 <clio-agent@sisyphuslabs.ai>
This commit is contained in:
Alexander
2026-09-13 20:01:22 +02:00
parent 7cac5ac870
commit 1d3e4b7e8a
4 changed files with 137 additions and 4 deletions
+46 -3
View File
@@ -1,23 +1,31 @@
import pytest import pytest
from mock import Mock, patch, call from mock import Mock, patch
from thefuck.entrypoints.fix_command import fix_command from thefuck.entrypoints.fix_command import fix_command
from thefuck.types import CorrectedCommand from thefuck.types import CorrectedCommand
@pytest.fixture @pytest.fixture
def mock_learned(monkeypatch): def mock_learned(monkeypatch):
state = {"correction": None, "recordings": []} state = {"correction": None, "guess": None, "recordings": []}
def fake_get_correction(script): def fake_get_correction(script):
return state["correction"] return state["correction"]
def fake_guess(script):
return state["guess"]
def fake_record(original, corrected): def fake_record(original, corrected):
state["recordings"].append((original, corrected)) state["recordings"].append((original, corrected))
monkeypatch.setattr( monkeypatch.setattr(
"thefuck.entrypoints.fix_command.get_correction", fake_get_correction "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 return state
@@ -75,6 +83,41 @@ class TestLearnedAutoApply(object):
assert shown_cmd.script == "git push origin main" 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): class TestRecordOnSelection(object):
def test_records_user_selection( def test_records_user_selection(
self, mock_learned, known_args, settings, monkeypatch self, mock_learned, known_args, settings, monkeypatch
+52
View File
@@ -108,6 +108,58 @@ class TestGetCorrection(object):
assert learned.get_correction("sl") == "ls" 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): class TestClear(object):
def test_removes_all_entries(self, learned): def test_removes_all_entries(self, learned):
learned.record("git psuh", "git push") learned.record("git psuh", "git push")
+7 -1
View File
@@ -6,7 +6,7 @@ from .. import logs, types, const
from ..conf import settings from ..conf import settings
from ..corrector import get_corrected_commands from ..corrector import get_corrected_commands
from ..exceptions import EmptyCommand 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 ..ui import select_command
from ..utils import get_alias, get_all_executables from ..utils import get_alias, get_all_executables
@@ -41,6 +41,12 @@ def fix_command(known_args):
return return
learned_script = get_correction(command.script) 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: if learned_script:
learned_cmd = types.CorrectedCommand( learned_cmd = types.CorrectedCommand(
script=learned_script, side_effect=None, priority=0 script=learned_script, side_effect=None, priority=0
+32
View File
@@ -2,8 +2,10 @@ import atexit
import os import os
import shelve import shelve
import time import time
from difflib import get_close_matches
from . import logs from . import logs
from .utils import get_all_executables, which
try: try:
import dbm import dbm
@@ -18,6 +20,9 @@ except ImportError:
_shelve_open_error = () _shelve_open_error = ()
GUESS_CUTOFF = 0.8
class LearnedCorrections(object): class LearnedCorrections(object):
def __init__(self): def __init__(self):
self._db = None self._db = None
@@ -137,6 +142,32 @@ class LearnedCorrections(object):
return None 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): def clear(self):
db = self.db db = self.db
for key in list(db.keys()): for key in list(db.keys()):
@@ -154,4 +185,5 @@ _learned = LearnedCorrections()
record = _learned.record record = _learned.record
get_correction = _learned.get_correction get_correction = _learned.get_correction
guess_from_path = _learned.guess_from_path
clear = _learned.clear clear = _learned.clear