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').
This commit is contained in:
@@ -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()
|
||||
@@ -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"
|
||||
@@ -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)
|
||||
|
||||
@@ -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
|
||||
Reference in New Issue
Block a user