diff --git a/tests/test_cli/test_app.py b/tests/test_cli/test_app.py new file mode 100644 index 0000000..3b91de6 --- /dev/null +++ b/tests/test_cli/test_app.py @@ -0,0 +1,23 @@ +from __future__ import annotations + +from pathlib import Path + +from click.testing import CliRunner + +from plyngent.cli.app import main + + +def test_providers_list() -> None: + config = Path(__file__).resolve().parents[1] / "test_config" / "plyngent-valid.toml" + runner = CliRunner() + result = runner.invoke(main, ["providers", "--config", str(config)]) + assert result.exit_code == 0 + assert "test1" in result.output + assert "openai" in result.output + + +def test_help() -> None: + runner = CliRunner() + result = runner.invoke(main, ["--help"]) + assert result.exit_code == 0 + assert "chat" in result.output diff --git a/tests/test_cli/test_repl_commands.py b/tests/test_cli/test_repl_commands.py new file mode 100644 index 0000000..b853c91 --- /dev/null +++ b/tests/test_cli/test_repl_commands.py @@ -0,0 +1,123 @@ +from __future__ import annotations + +from typing import TYPE_CHECKING, Literal, overload + +import pytest +import tomlkit + +from plyngent.agent import ChatAgent +from plyngent.cli.repl import handle_slash +from plyngent.cli.state import ReplState +from plyngent.config.models import DatabaseConfig, OpenAIProvider +from plyngent.config.store import ConfigStore +from plyngent.lmproto.openai_compatible.model import ( + AssistantChatMessage, + ChatCompletionChoice, + ChatCompletionChunk, + ChatCompletionResponse, + ChatCompletionsParam, +) +from plyngent.memory import MemoryStore +from plyngent.tools import set_workspace_root + +if TYPE_CHECKING: + from collections.abc import AsyncIterator + from pathlib import Path + + +class DummyClient: + @overload + async def chat_completions( + self, param: ChatCompletionsParam, *, stream: Literal[False] = False + ) -> ChatCompletionResponse: ... + + @overload + async def chat_completions( + self, param: ChatCompletionsParam, *, stream: Literal[True] + ) -> AsyncIterator[ChatCompletionChunk]: ... + + async def chat_completions( + self, param: ChatCompletionsParam, *, stream: bool = False + ) -> ChatCompletionResponse | AsyncIterator[ChatCompletionChunk]: + del param + if stream: + + async def empty() -> AsyncIterator[ChatCompletionChunk]: + empty_chunks: list[ChatCompletionChunk] = [] + for chunk in empty_chunks: + yield chunk + + return empty() + return ChatCompletionResponse( + id="1", + object="chat.completion", + created=0, + model="m", + choices=[ + ChatCompletionChoice( + index=0, + message=AssistantChatMessage(content="ok"), + logprobs={}, + finish_reason="stop", + ) + ], + system_fingerprint="", + usage={}, + ) + + +@pytest.fixture +async def state(tmp_path: Path) -> AsyncIterator[ReplState]: + _ = set_workspace_root(tmp_path) + memory = await MemoryStore.open(DatabaseConfig()) + provider = OpenAIProvider(access_key_or_token="sk-test") + config = ConfigStore(path=tmp_path / "plyngent.toml", document=tomlkit.document()) + config.providers = {"local": provider} + st = ReplState( + config=config, + memory=memory, + workspace=tmp_path, + provider_name="local", + provider=provider, + model="gpt-test", + tools_enabled=False, + ) + st.client = DummyClient() + st.agent = ChatAgent(st.client, model=st.model, memory=st.memory, session_id=None) + await st.new_session("t") + yield st + await memory.close() + + +async def test_help_and_clear(state: ReplState) -> None: + assert await handle_slash(state, "/help") is True + state.agent.messages.append(AssistantChatMessage(content="x")) + assert await handle_slash(state, "/clear") is True + assert state.agent.messages == [] + + +async def test_quit(state: ReplState) -> None: + assert await handle_slash(state, "/quit") is False + + +async def test_new_and_sessions(state: ReplState, capsys: pytest.CaptureFixture[str]) -> None: + first = state.session_id + assert await handle_slash(state, "/new other") is True + assert state.session_id != first + assert await handle_slash(state, "/sessions") is True + out = capsys.readouterr().out + assert str(state.session_id) in out + + +async def test_tools_toggle(state: ReplState) -> None: + assert await handle_slash(state, "/tools on") is True + assert state.tools_enabled is True + assert await handle_slash(state, "/tools off") is True + assert state.tools_enabled is False + + +async def test_resume(state: ReplState) -> None: + sid = state.session_id + assert sid is not None + state.agent.messages.clear() + assert await handle_slash(state, f"/resume {sid}") is True diff --git a/tests/test_cli/test_selection.py b/tests/test_cli/test_selection.py new file mode 100644 index 0000000..2fe85aa --- /dev/null +++ b/tests/test_cli/test_selection.py @@ -0,0 +1,48 @@ +from __future__ import annotations + +import pytest + +from plyngent.cli.selection import select_model, select_provider +from plyngent.config.models import ModelConfig, OpenAICompatibleProvider, OpenAIProvider + + +def test_select_provider_preferred() -> None: + providers = { + "a": OpenAIProvider(access_key_or_token="sk"), + "b": OpenAICompatibleProvider(access_key_or_token="sk", url="https://x/v1"), + } + name, provider = select_provider(providers, preferred="b") + assert name == "b" + assert isinstance(provider, OpenAICompatibleProvider) + + +def test_select_provider_single_auto() -> None: + providers = {"only": OpenAIProvider(access_key_or_token="sk")} + name, _ = select_provider(providers) + assert name == "only" + + +def test_select_provider_unknown() -> None: + providers = {"a": OpenAIProvider(access_key_or_token="sk")} + with pytest.raises(Exception, match="unknown provider"): + _ = select_provider(providers, preferred="nope") + + +def test_select_model_from_list() -> None: + provider = OpenAICompatibleProvider( + access_key_or_token="sk", + url="https://x/v1", + models={"m1": ModelConfig()}, + ) + assert select_model(provider) == "m1" + assert select_model(provider, preferred="m1") == "m1" + + +def test_select_model_prompt_when_empty(monkeypatch: pytest.MonkeyPatch) -> None: + provider = OpenAIProvider(access_key_or_token="sk") + + def _prompt(*_args: object, **_kwargs: object) -> str: + return "gpt-test" + + monkeypatch.setattr("click.prompt", _prompt) + assert select_model(provider) == "gpt-test"