diff --git a/tests/entrypoints/test_fix_command_learned.py b/tests/entrypoints/test_fix_command_learned.py index 26b1a02..2a8dd07 100644 --- a/tests/entrypoints/test_fix_command_learned.py +++ b/tests/entrypoints/test_fix_command_learned.py @@ -6,26 +6,41 @@ from thefuck.types import CorrectedCommand @pytest.fixture def mock_learned(monkeypatch): - state = {"correction": None, "guess": None, "recordings": []} + state = {"correction": None, "guess": None, "history": None, + "help": None, "recordings": [], "calls": []} def fake_get_correction(script): + state["calls"].append("correction") return state["correction"] + def fake_history(command): + state["calls"].append("history") + return state["history"] + def fake_guess(script): + state["calls"].append("guess") return state["guess"] + def fake_help(script): + state["calls"].append("help") + return state["help"] + 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.get_history_correction", fake_history + ) monkeypatch.setattr( "thefuck.entrypoints.fix_command.guess_from_path", fake_guess ) monkeypatch.setattr( - "thefuck.entrypoints.fix_command.record", fake_record + "thefuck.entrypoints.fix_command.get_help_correction", fake_help ) + monkeypatch.setattr("thefuck.entrypoints.fix_command.record", fake_record) return state @@ -118,6 +133,87 @@ class TestGuessAutoApply(object): assert mock_learned["recordings"] == [] +class TestResolverChainOrder(object): + def test_full_chain_consulted_in_order_when_only_help_hits( + self, mock_learned, known_args, settings, monkeypatch + ): + mock_learned["help"] = "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) + + assert mock_learned["calls"] == [ + "correction", "history", "guess", "help" + ] + get_corrected.assert_not_called() + + @pytest.mark.parametrize("hit_source,expected_calls", [ + ("history", ["correction", "history"]), + ("guess", ["correction", "history", "guess"]), + ]) + def test_earlier_hit_stops_the_chain( + self, mock_learned, known_args, settings, hit_source, expected_calls + ): + mock_learned[hit_source] = "git push origin main" + + with patch("thefuck.types.CorrectedCommand.run"), patch( + "thefuck.logs.show_corrected_command" + ): + fix_command(known_args) + + assert mock_learned["calls"] == expected_calls + + def test_history_receives_command_object( + self, mock_learned, known_args, settings, monkeypatch + ): + history_mock = Mock(return_value=None) + monkeypatch.setattr( + "thefuck.entrypoints.fix_command.get_history_correction", history_mock + ) + 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 history_mock.call_count == 1 + assert history_mock.call_args[0][0].script == "git psuh origin main" + + +class TestResolverAutoApply(object): + @pytest.mark.parametrize("source", ["history", "help"]) + def test_resolver_hit_records_and_auto_applies( + self, mock_learned, known_args, settings, monkeypatch, source + ): + mock_learned[source] = "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" + ) as mock_show: + fix_command(known_args) + mock_run.assert_called_once() + get_corrected.assert_not_called() + assert mock_show.call_args[0][0].script == "git push origin main" + assert mock_learned["recordings"] == [ + ("git psuh origin main", "git push origin main") + ] + + class TestRecordOnSelection(object): def test_records_user_selection( self, mock_learned, known_args, settings, monkeypatch @@ -172,3 +268,27 @@ class TestFallthrough(object): with patch("thefuck.types.CorrectedCommand.run") as mock_run: fix_command(known_args) mock_run.assert_called_once() + + def test_falls_through_to_rules_when_all_resolvers_decline( + self, mock_learned, known_args, settings, monkeypatch + ): + selected = CorrectedCommand( + script="git push origin main", side_effect=None, priority=100 + ) + select = Mock(return_value=selected) + monkeypatch.setattr( + "thefuck.entrypoints.fix_command.get_corrected_commands", + lambda _: iter([selected]), + ) + monkeypatch.setattr( + "thefuck.entrypoints.fix_command.select_command", select + ) + + with patch("thefuck.types.CorrectedCommand.run") as mock_run: + fix_command(known_args) + mock_run.assert_called_once() + + assert mock_learned["calls"] == [ + "correction", "history", "guess", "help" + ] + assert select.call_count == 1 diff --git a/thefuck/entrypoints/fix_command.py b/thefuck/entrypoints/fix_command.py index 22bdd15..cd63533 100644 --- a/thefuck/entrypoints/fix_command.py +++ b/thefuck/entrypoints/fix_command.py @@ -7,6 +7,8 @@ from ..conf import settings from ..corrector import get_corrected_commands from ..exceptions import EmptyCommand from ..learned import get_correction, guess_from_path, record +from ..resolvers.help_resolver import get_help_correction +from ..resolvers.history_resolver import get_history_correction from ..ui import select_command from ..utils import get_alias, get_all_executables @@ -41,12 +43,24 @@ def fix_command(known_args): return learned_script = get_correction(command.script) + if not learned_script: + learned_script = get_history_correction(command) + if learned_script: + logs.debug("Corrected from history: {}".format( + learned_script)) + record(command.script, learned_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 not learned_script: + learned_script = get_help_correction(command.script) + if learned_script: + logs.debug("Corrected from binary help: {}".format( + learned_script)) + record(command.script, learned_script) if learned_script: learned_cmd = types.CorrectedCommand( script=learned_script, side_effect=None, priority=0