From 8d2f7d13c7365289aaa3f0cc1ebacd76e1b98889 Mon Sep 17 00:00:00 2001 From: Alexander Date: Wed, 15 Apr 2026 12:06:41 +0200 Subject: [PATCH] Add dynamic learning from accepted corrections Record user-accepted corrections in a shelve-backed store and replay them on future matching mistakes, bypassing rule evaluation entirely. Two-level matching: exact full-command lookup, then word-level reconstruction that generalises to different arguments (e.g. learning 'git psuh origin main' also fixes 'git psuh origin dev'). --- tests/entrypoints/test_fix_command_learned.py | 131 +++++++++++++++ tests/test_learned.py | 139 ++++++++++++++++ thefuck/entrypoints/fix_command.py | 21 ++- thefuck/learned.py | 157 ++++++++++++++++++ 4 files changed, 443 insertions(+), 5 deletions(-) create mode 100644 tests/entrypoints/test_fix_command_learned.py create mode 100644 tests/test_learned.py create mode 100644 thefuck/learned.py diff --git a/tests/entrypoints/test_fix_command_learned.py b/tests/entrypoints/test_fix_command_learned.py new file mode 100644 index 0000000..a5f22e1 --- /dev/null +++ b/tests/entrypoints/test_fix_command_learned.py @@ -0,0 +1,131 @@ +import pytest +from mock import Mock, patch, call +from thefuck.entrypoints.fix_command import fix_command +from thefuck.types import CorrectedCommand + + +@pytest.fixture +def mock_learned(monkeypatch): + state = {"correction": None, "recordings": []} + + def fake_get_correction(script): + return state["correction"] + + 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) + return state + + +@pytest.fixture +def known_args(): + return Mock( + force_command="git psuh origin main", yes=False, debug=False, repeat=False + ) + + +class TestLearnedAutoApply(object): + def test_auto_applies_learned_correction( + self, mock_learned, known_args, settings, monkeypatch + ): + mock_learned["correction"] = "git push origin main" + monkeypatch.setattr( + "thefuck.entrypoints.fix_command.get_corrected_commands", lambda _: iter([]) + ) + monkeypatch.setattr( + "thefuck.entrypoints.fix_command.select_command", lambda _: None + ) + + with patch("thefuck.types.CorrectedCommand.run") as mock_run, patch( + "thefuck.logs.show_corrected_command" + ): + fix_command(known_args) + mock_run.assert_called_once() + + def test_learned_skips_rule_matching( + self, mock_learned, known_args, settings, monkeypatch + ): + mock_learned["correction"] = "git push origin main" + get_corrected = Mock() + monkeypatch.setattr( + "thefuck.entrypoints.fix_command.get_corrected_commands", get_corrected + ) + + with patch("thefuck.types.CorrectedCommand.run"), patch( + "thefuck.logs.show_corrected_command" + ): + fix_command(known_args) + get_corrected.assert_not_called() + + def test_shows_corrected_command_on_auto_apply( + self, mock_learned, known_args, settings, monkeypatch + ): + mock_learned["correction"] = "git push origin main" + + with patch("thefuck.types.CorrectedCommand.run"), patch( + "thefuck.logs.show_corrected_command" + ) as mock_show: + fix_command(known_args) + assert mock_show.call_count == 1 + shown_cmd = mock_show.call_args[0][0] + assert shown_cmd.script == "git push origin main" + + +class TestRecordOnSelection(object): + def test_records_user_selection( + self, mock_learned, known_args, settings, monkeypatch + ): + selected = CorrectedCommand( + script="git push origin main", side_effect=None, priority=100 + ) + monkeypatch.setattr( + "thefuck.entrypoints.fix_command.get_corrected_commands", + lambda _: iter([selected]), + ) + monkeypatch.setattr( + "thefuck.entrypoints.fix_command.select_command", lambda _: selected + ) + + with patch("thefuck.types.CorrectedCommand.run"): + fix_command(known_args) + assert mock_learned["recordings"] == [ + ("git psuh origin main", "git push origin main") + ] + + def test_does_not_record_on_abort( + self, mock_learned, known_args, settings, monkeypatch + ): + monkeypatch.setattr( + "thefuck.entrypoints.fix_command.get_corrected_commands", lambda _: iter([]) + ) + monkeypatch.setattr( + "thefuck.entrypoints.fix_command.select_command", lambda _: None + ) + + with pytest.raises(SystemExit): + fix_command(known_args) + assert mock_learned["recordings"] == [] + + +class TestFallthrough(object): + def test_falls_through_to_rules_when_no_learned( + self, mock_learned, known_args, settings, monkeypatch + ): + selected = CorrectedCommand( + script="git push origin main", side_effect=None, priority=100 + ) + monkeypatch.setattr( + "thefuck.entrypoints.fix_command.get_corrected_commands", + lambda _: iter([selected]), + ) + monkeypatch.setattr( + "thefuck.entrypoints.fix_command.select_command", lambda _: selected + ) + + with patch("thefuck.types.CorrectedCommand.run") as mock_run: + fix_command(known_args) + mock_run.assert_called_once() diff --git a/tests/test_learned.py b/tests/test_learned.py new file mode 100644 index 0000000..c9bab33 --- /dev/null +++ b/tests/test_learned.py @@ -0,0 +1,139 @@ +import pytest +from thefuck.learned import LearnedCorrections + + +@pytest.fixture +def learned(tmp_path): + lc = LearnedCorrections() + lc._db = {} + return lc + + +class TestRecord(object): + def test_noop_when_scripts_identical(self, learned): + learned.record("git push", "git push") + assert len(learned.db) == 0 + + def test_stores_full_command_mapping(self, learned): + learned.record("git psuh origin main", "git push origin main") + assert ( + learned.db["cmd:git psuh origin main"]["corrected"] + == "git push origin main" + ) + + def test_increments_count_on_repeat(self, learned): + learned.record("git psuh", "git push") + learned.record("git psuh", "git push") + assert learned.db["cmd:git psuh"]["count"] == 2 + + def test_updates_timestamp(self, learned): + learned.record("git psuh", "git push") + first_ts = learned.db["cmd:git psuh"]["timestamp"] + learned.record("git psuh", "git push") + assert learned.db["cmd:git psuh"]["timestamp"] >= first_ts + + def test_stores_command_word_replacement(self, learned): + learned.record("pyhton script.py", "python script.py") + assert learned.db["word:pyhton"]["replacement"] == "python" + + def test_stores_subcommand_replacement(self, learned): + learned.record("git psuh origin main", "git push origin main") + assert learned.db["part:git:psuh"]["replacement"] == "push" + + def test_stores_multiple_word_diffs(self, learned): + learned.record("gti comit -m msg", "git commit -m msg") + assert learned.db["word:gti"]["replacement"] == "git" + assert learned.db["part:git:comit"]["replacement"] == "commit" + + def test_no_word_level_when_lengths_differ(self, learned): + learned.record("git push", "git push --set-upstream origin main") + assert "cmd:git push" in learned.db + assert not any( + k.startswith("word:") or k.startswith("part:") for k in learned.db + ) + + def test_updates_correction_on_changed_choice(self, learned): + learned.record("apt install vim", "sudo apt install vim") + learned.record("apt install vim", "apt-get install vim") + assert learned.db["cmd:apt install vim"]["corrected"] == "apt-get install vim" + assert learned.db["cmd:apt install vim"]["count"] == 2 + + +class TestGetCorrection(object): + def test_exact_full_command_match(self, learned): + learned.record("git psuh origin main", "git push origin main") + assert learned.get_correction("git psuh origin main") == "git push origin main" + + def test_generalises_subcommand_to_different_args(self, learned): + learned.record("git psuh origin main", "git push origin main") + assert learned.get_correction("git psuh origin dev") == "git push origin dev" + + def test_generalises_command_word(self, learned): + learned.record("pyhton script.py", "python script.py") + assert learned.get_correction("pyhton other.py") == "python other.py" + + def test_combined_command_and_subcommand(self, learned): + learned.record("gti comit -m msg", "git commit -m msg") + assert learned.get_correction('gti comit -m "other"') == 'git commit -m "other"' + + def test_returns_none_when_no_match(self, learned): + assert learned.get_correction("totally unknown cmd") is None + + def test_returns_none_for_empty_script(self, learned): + assert learned.get_correction("") is None + + def test_prefers_full_command_over_word_level(self, learned): + learned.record("git psuh origin main", "git push origin main") + learned.db["cmd:git psuh origin main"]["corrected"] = ( + "git push --force origin main" + ) + assert ( + learned.get_correction("git psuh origin main") + == "git push --force origin main" + ) + + def test_word_level_only_replaces_known_tokens(self, learned): + learned.record("git psuh origin main", "git push origin main") + result = learned.get_correction("git psuh origin dev") + assert result == "git push origin dev" + + def test_cross_resolve_command_and_part(self, learned): + """When both cmd word and subcommand are typos, parts stored + under the corrected cmd name still resolve.""" + learned.record("gti psuh origin main", "git push origin main") + assert learned.get_correction("gti psuh origin dev") == "git push origin dev" + + def test_single_word_command(self, learned): + learned.record("sl", "ls") + assert learned.get_correction("sl") == "ls" + + +class TestClear(object): + def test_removes_all_entries(self, learned): + learned.record("git psuh", "git push") + learned.record("pyhton x.py", "python x.py") + learned.clear() + assert len(learned.db) == 0 + + def test_no_matches_after_clear(self, learned): + learned.record("git psuh", "git push") + learned.clear() + assert learned.get_correction("git psuh") is None + + +class TestRoundTrip(object): + def test_record_then_match(self, learned): + learned.record("docker bilud .", "docker build .") + assert learned.get_correction("docker bilud .") == "docker build ." + + def test_record_then_generalise(self, learned): + learned.record("docker bilud -t foo .", "docker build -t foo .") + assert ( + learned.get_correction("docker bilud -t bar .") == "docker build -t bar ." + ) + + def test_multiple_distinct_commands(self, learned): + learned.record("git psuh", "git push") + learned.record("pyhton x.py", "python x.py") + assert learned.get_correction("git psuh") == "git push" + assert learned.get_correction("pyhton y.py") == "python y.py" diff --git a/thefuck/entrypoints/fix_command.py b/thefuck/entrypoints/fix_command.py index 018ba58..9722aee 100644 --- a/thefuck/entrypoints/fix_command.py +++ b/thefuck/entrypoints/fix_command.py @@ -6,6 +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 ..ui import select_command from ..utils import get_alias, get_all_executables @@ -13,10 +14,10 @@ from ..utils import get_alias, get_all_executables def _get_raw_command(known_args): if known_args.force_command: return [known_args.force_command] - elif not os.environ.get('TF_HISTORY'): + elif not os.environ.get("TF_HISTORY"): return known_args.command else: - history = os.environ['TF_HISTORY'].split('\n')[::-1] + history = os.environ["TF_HISTORY"].split("\n")[::-1] alias = get_alias() executables = get_all_executables() for command in history: @@ -29,20 +30,30 @@ def _get_raw_command(known_args): def fix_command(known_args): """Fixes previous command. Used when `thefuck` called without arguments.""" settings.init(known_args) - with logs.debug_time('Total'): - logs.debug(u'Run with settings: {}'.format(pformat(settings))) + with logs.debug_time("Total"): + logs.debug("Run with settings: {}".format(pformat(settings))) raw_command = _get_raw_command(known_args) try: command = types.Command.from_raw_script(raw_command) except EmptyCommand: - logs.debug('Empty command, nothing to do') + logs.debug("Empty command, nothing to do") + return + + learned_script = get_correction(command.script) + if learned_script: + learned_cmd = types.CorrectedCommand( + script=learned_script, side_effect=None, priority=0 + ) + logs.show_corrected_command(learned_cmd) + learned_cmd.run(command) return corrected_commands = get_corrected_commands(command) selected_command = select_command(corrected_commands) if selected_command: + record(command.script, selected_command.script) selected_command.run(command) else: sys.exit(1) diff --git a/thefuck/learned.py b/thefuck/learned.py new file mode 100644 index 0000000..ec74b13 --- /dev/null +++ b/thefuck/learned.py @@ -0,0 +1,157 @@ +import atexit +import os +import shelve +import time + +from . import logs + +try: + import dbm + + _shelve_open_error = (dbm.error,) +except ImportError: + try: + import anydbm + + _shelve_open_error = (anydbm.error,) + except ImportError: + _shelve_open_error = () + + +class LearnedCorrections(object): + def __init__(self): + self._db = None + + def _init_db(self): + try: + self._setup_db() + except Exception: + logs.debug("Unable to init learned-corrections db") + self._db = {} + + def _setup_db(self): + cache_dir = self._get_cache_dir() + cache_path = os.path.join(cache_dir, "thefuck_learned") + try: + self._db = shelve.open(cache_path) + except _shelve_open_error + (ImportError,): + logs.warn("Removing possibly out-dated learned-corrections db") + for suffix in ("", ".db", ".dir", ".bak", ".dat"): + path = cache_path + suffix + if os.path.exists(path): + os.remove(path) + self._db = shelve.open(cache_path) + atexit.register(self._db.close) + + @staticmethod + def _get_cache_dir(): + cache_dir = os.getenv("XDG_CACHE_HOME", os.path.expanduser("~/.cache")) + try: + os.makedirs(cache_dir) + except OSError: + if not os.path.isdir(cache_dir): + raise + return cache_dir + + @property + def db(self): + if self._db is None: + self._init_db() + return self._db + + def record(self, original_script, corrected_script): + if original_script == corrected_script: + return + + db = self.db + now = time.time() + + original_parts = original_script.split() + corrected_parts = corrected_script.split() + + full_key = "cmd:" + original_script + entry = db.get(full_key, {}) + entry["corrected"] = corrected_script + entry["count"] = entry.get("count", 0) + 1 + entry["timestamp"] = now + db[full_key] = entry + + # Word-level diffs: only when token counts match, store each + # changed token keyed by position so lookups can generalise + # (e.g. learning "git psuh origin main" also fixes "git psuh origin dev") + if ( + original_parts + and corrected_parts + and len(original_parts) == len(corrected_parts) + ): + for i, (orig_tok, corr_tok) in enumerate( + zip(original_parts, corrected_parts) + ): + if orig_tok == corr_tok: + continue + if i == 0: + key = "word:" + orig_tok + else: + # Keyed under the corrected cmd name so "gti psuh" + # resolves via word:gti→git then part:git:psuh→push + key = "part:" + corrected_parts[0] + ":" + orig_tok + part_entry = db.get(key, {}) + part_entry["replacement"] = corr_tok + part_entry["count"] = part_entry.get("count", 0) + 1 + part_entry["timestamp"] = now + db[key] = part_entry + + self._sync() + + def get_correction(self, script): + db = self.db + + full_key = "cmd:" + script + entry = db.get(full_key) + if entry: + return entry["corrected"] + + parts = script.split() + if not parts: + return None + + corrected_parts = list(parts) + found = False + + word_entry = db.get("word:" + parts[0]) + if word_entry: + corrected_parts[0] = word_entry["replacement"] + found = True + + # Part lookups use the (possibly corrected) cmd name so that + # "gti psuh" resolves even though parts are stored under "git" + cmd_name = corrected_parts[0] + for i in range(1, len(parts)): + part_entry = db.get("part:" + cmd_name + ":" + parts[i]) + if part_entry: + corrected_parts[i] = part_entry["replacement"] + found = True + + if found: + return " ".join(corrected_parts) + + return None + + def clear(self): + db = self.db + for key in list(db.keys()): + del db[key] + self._sync() + + def _sync(self): + try: + self.db.sync() + except AttributeError: + pass + + +_learned = LearnedCorrections() + +record = _learned.record +get_correction = _learned.get_correction +clear = _learned.clear