diff --git a/llm/default_plugins/default_tools.py b/llm/default_plugins/default_tools.py index 53ff72cd0..9eedc07fc 100644 --- a/llm/default_plugins/default_tools.py +++ b/llm/default_plugins/default_tools.py @@ -1,8 +1,9 @@ import llm -from llm.tools import llm_time, llm_version +from llm.tools import llm_random_choice, llm_time, llm_version @llm.hookimpl def register_tools(register): register(llm_version) register(llm_time) + register(llm_random_choice) diff --git a/llm/tools.py b/llm/tools.py index 0970205c7..405925af3 100644 --- a/llm/tools.py +++ b/llm/tools.py @@ -1,3 +1,4 @@ +import random import time from datetime import datetime, timezone from importlib.metadata import version @@ -35,3 +36,8 @@ def llm_time() -> dict: "timezone_offset": timezone_offset, "is_dst": is_dst, } + + +def llm_random_choice(choices: list[str]) -> str: + "Return a random choice from a list of strings" + return random.choice(choices) diff --git a/tests/test_tools.py b/tests/test_tools.py index 93c95ba39..8a1b8ce5d 100644 --- a/tests/test_tools.py +++ b/tests/test_tools.py @@ -12,7 +12,7 @@ import llm from llm import CancelToolCall, cli from llm.migrations import migrate -from llm.tools import llm_time +from llm.tools import llm_random_choice, llm_time API_KEY = os.environ.get("PYTEST_OPENAI_API_KEY", None) or "badkey" @@ -380,6 +380,33 @@ def test_default_tool_llm_time(): } +def test_default_tool_llm_random_choice(monkeypatch): + monkeypatch.setattr("llm.tools.random.choice", lambda choices: choices[-1]) + runner = CliRunner() + result = runner.invoke( + cli.cli, + [ + "-m", + "echo", + "-T", + "llm_random_choice", + json.dumps( + { + "tool_calls": [ + { + "name": "llm_random_choice", + "arguments": {"choices": ["red", "green", "blue"]}, + } + ] + } + ), + ], + ) + assert result.exit_code == 0 + assert '"output": "blue"' in result.output + assert llm_random_choice(["one", "two"]) == "two" + + def test_incorrect_tool_usage(): model = llm.get_model("echo")