diff --git a/src/goose/cli/main.py b/src/goose/cli/main.py index a0495a30..e60facc3 100644 --- a/src/goose/cli/main.py +++ b/src/goose/cli/main.py @@ -97,11 +97,24 @@ def list_toolkits() -> None: print(f" - [bold]{toolkit_name}[/bold]: {first_line_of_doc}") +def autocomplete_session_files(ctx: click.Context, args: str, incomplete: str) -> None: + return [ + f"{session_name}" + for session_name in sorted(get_session_files().keys(), reverse=True, key=lambda x: x.lower()) + if session_name.startswith(incomplete) + ] + + +def get_session_files() -> dict[str, Path]: + return list_sorted_session_files(SESSIONS_PATH) + + @session.command(name="start") +@click.argument("name", required=False, shell_complete=autocomplete_session_files) @click.option("--profile") @click.option("--plan", type=click.Path(exists=True)) @click.option("--log-level", type=click.Choice(["DEBUG", "INFO", "WARNING", "ERROR", "CRITICAL"]), default="INFO") -def session_start(profile: str, log_level: str, plan: Optional[str] = None) -> None: +def session_start(name: Optional[str], profile: str, log_level: str, plan: Optional[str] = None) -> None: """Start a new goose session""" if plan: yaml = YAML() @@ -109,7 +122,7 @@ def session_start(profile: str, log_level: str, plan: Optional[str] = None) -> N _plan = yaml.load(f) else: _plan = None - session = Session(profile=profile, plan=_plan, log_level=log_level) + session = Session(name=name, profile=profile, plan=_plan, log_level=log_level) session.run() @@ -126,30 +139,20 @@ def parse_args(ctx: click.Context, param: click.Parameter, value: str) -> dict[s @session.command(name="planned") @click.option("--plan", type=click.Path(exists=True)) +@click.option("--log-level", type=click.Choice(["DEBUG", "INFO", "WARNING", "ERROR", "CRITICAL"]), default="INFO") @click.option("-a", "--args", callback=parse_args, help="Args in the format arg1:value1,arg2:value2") -def session_planned(plan: str, args: Optional[dict[str, str]]) -> None: +def session_planned(plan: str, log_level: str, args: Optional[dict[str, str]]) -> None: plan_templated = render_template(Path(plan), context=args) _plan = parse_plan(plan_templated) - session = Session(plan=_plan) + session = Session(plan=_plan, log_level=log_level) session.run() -def autocomplete_session_files(ctx: click.Context, args: str, incomplete: str) -> None: - return [ - f"{session_name}" - for session_name in sorted(get_session_files().keys(), reverse=True, key=lambda x: x.lower()) - if session_name.startswith(incomplete) - ] - - -def get_session_files() -> dict[str, Path]: - return list_sorted_session_files(SESSIONS_PATH) - - @session.command(name="resume") @click.argument("name", required=False, shell_complete=autocomplete_session_files) @click.option("--profile") -def session_resume(name: Optional[str], profile: str) -> None: +@click.option("--log-level", type=click.Choice(["DEBUG", "INFO", "WARNING", "ERROR", "CRITICAL"]), default="INFO") +def session_resume(name: Optional[str], profile: str, log_level: str) -> None: """Resume an existing goose session""" session_files = get_session_files() if name is None: @@ -164,7 +167,7 @@ def session_resume(name: Optional[str], profile: str) -> None: print(f"Resuming session: {name}") else: print(f"Creating new session: {name}") - session = Session(name=name, profile=profile) + session = Session(name=name, profile=profile, log_level=log_level) session.run() diff --git a/src/goose/cli/session.py b/src/goose/cli/session.py index 5674eb48..38128d85 100644 --- a/src/goose/cli/session.py +++ b/src/goose/cli/session.py @@ -3,7 +3,6 @@ from pathlib import Path from typing import Any, Dict, List, Optional from exchange import Message, ToolResult, ToolUse, Text -from prompt_toolkit.shortcuts import confirm from rich import print from rich.console import RenderableType from rich.live import Live @@ -19,7 +18,7 @@ from goose.notifier import Notifier from goose.profile import Profile from goose.utils import droid, load_plugins from goose.utils._cost_calculator import get_total_cost_message -from goose.utils.session_file import read_from_file, write_to_file +from goose.utils.session_file import read_or_create_file, save_latest_session RESUME_MESSAGE = "I see we were interrupted. How can I help you?" @@ -84,38 +83,44 @@ class Session: log_level: Optional[str] = "INFO", **kwargs: Dict[str, Any], ) -> None: - self.name = name + if name is None: + self.name = droid() + print(Panel(f"Session name not provided, using generated name: {self.name}")) + else: + self.name = name self.status_indicator = Status("", spinner="dots") self.notifier = SessionNotifier(self.status_indicator) self.exchange = build_exchange(profile=load_profile(profile), notifier=self.notifier) setup_logging(log_file_directory=LOG_PATH, log_level=log_level) - if name is not None and self.session_file_path.exists(): - messages = self.load_session() - - if messages and messages[-1].role == "user": - if type(messages[-1].content[-1]) is Text: - # remove the last user message - messages.pop() - elif type(messages[-1].content[-1]) is ToolResult: - # if we remove this message, we would need to remove - # the previous assistant message as well. instead of doing - # that, we just add a new assistant message to prompt the user - messages.append(Message.assistant(RESUME_MESSAGE)) - if messages and type(messages[-1].content[-1]) is ToolUse: - # remove the last request for a tool to be used - messages.pop() - - # add a new assistant text message to prompt the user - messages.append(Message.assistant(RESUME_MESSAGE)) - self.exchange.messages.extend(messages) + self.exchange.messages.extend(self._get_initial_messages()) if len(self.exchange.messages) == 0 and plan: self.setup_plan(plan=plan) self.prompt_session = GoosePromptSession() + def _get_initial_messages(self) -> List[Message]: + messages = self.load_session() + + if messages and messages[-1].role == "user": + if type(messages[-1].content[-1]) is Text: + # remove the last user message + messages.pop() + elif type(messages[-1].content[-1]) is ToolResult: + # if we remove this message, we would need to remove + # the previous assistant message as well. instead of doing + # that, we just add a new assistant message to prompt the user + messages.append(Message.assistant(RESUME_MESSAGE)) + if messages and type(messages[-1].content[-1]) is ToolUse: + # remove the last request for a tool to be used + messages.pop() + + # add a new assistant text message to prompt the user + messages.append(Message.assistant(RESUME_MESSAGE)) + return messages + def setup_plan(self, plan: dict) -> None: if len(self.exchange.messages): raise ValueError("The plan can only be set on an empty session.") @@ -160,12 +165,11 @@ class Session: + " - [yellow]depending on the error you may be able to continue[/]" ) self.notifier.stop() - + save_latest_session(self.session_file_path, self.exchange.messages) print() # Print a newline for separation. user_input = self.prompt_session.get_user_input() message = Message.user(text=user_input.text) if user_input.to_continue() else None - self.save_session() self._log_cost() def reply(self) -> None: @@ -226,29 +230,13 @@ class Session: def session_file_path(self) -> Path: return session_path(self.name) - def save_session(self) -> None: - """Save the current session to a file in JSON format.""" - if self.name is None: - self.generate_session_name() - - try: - if self.session_file_path.exists(): - if not confirm(f"Session {self.name} exists in {self.session_file_path}, overwrite?"): - self.generate_session_name() - write_to_file(self.session_file_path, self.exchange.messages) - except PermissionError as e: - raise RuntimeError(f"Failed to save session due to permissions: {e}") - except (IOError, OSError) as e: - raise RuntimeError(f"Failed to save session due to I/O error: {e}") - def load_session(self) -> List[Message]: - """Load a session from a JSON file.""" - return read_from_file(self.session_file_path) - - def generate_session_name(self) -> None: - user_entered_session_name = self.prompt_session.get_save_session_name() - self.name = user_entered_session_name if user_entered_session_name else droid() - print(f"Saving to [bold cyan]{self.session_file_path}[/bold cyan]") + message = ( + f"session is going to be saved to [bold cyan]{self.session_file_path}[/bold cyan]." + + " You can view it anytime." + ) + print(Panel(message)) + return read_or_create_file(self.session_file_path) def _log_cost(self) -> None: get_logger().info(get_total_cost_message(self.exchange.get_token_usage())) diff --git a/src/goose/utils/session_file.py b/src/goose/utils/session_file.py index a47efcb1..435186ce 100644 --- a/src/goose/utils/session_file.py +++ b/src/goose/utils/session_file.py @@ -1,5 +1,7 @@ import json +import os from pathlib import Path +import tempfile from typing import Dict, Iterator, List from exchange import Message @@ -9,9 +11,15 @@ from goose.cli.config import SESSION_FILE_SUFFIX def write_to_file(file_path: Path, messages: List[Message]) -> None: with open(file_path, "w") as f: - for m in messages: - json.dump(m.to_dict(), f) - f.write("\n") + _write_messages_to_file(f, messages) + + +def read_or_create_file(file_path: Path) -> List[Message]: + if file_path.exists(): + return read_from_file(file_path) + with open(file_path, "w"): + pass + return [] def read_from_file(file_path: Path) -> List[Message]: @@ -37,3 +45,17 @@ def session_file_exists(session_files_directory: Path) -> bool: if not session_files_directory.exists(): return False return any(list_session_files(session_files_directory)) + + +def save_latest_session(file_path: Path, messages: List[Message]) -> None: + with tempfile.NamedTemporaryFile("w", delete=False) as temp_file: + _write_messages_to_file(temp_file, messages) + temp_file_path = temp_file.name + + os.replace(temp_file_path, file_path) + + +def _write_messages_to_file(file: any, messages: List[Message]) -> None: + for m in messages: + json.dump(m.to_dict(), file) + file.write("\n") diff --git a/tests/cli/test_main.py b/tests/cli/test_main.py index b885f7f9..9f6f2f66 100644 --- a/tests/cli/test_main.py +++ b/tests/cli/test_main.py @@ -30,11 +30,19 @@ def mock_session(): yield mock_session_class, mock_session_instance +def test_session_start_command_with_session_name(mock_session): + mock_session_class, mock_session_instance = mock_session + runner = CliRunner() + runner.invoke(goose_cli, ["session", "start", "session1", "--profile", "default"]) + mock_session_class.assert_called_once_with(name="session1", profile="default", plan=None, log_level="INFO") + mock_session_instance.run.assert_called_once() + + def test_session_resume_command_with_session_name(mock_session): mock_session_class, mock_session_instance = mock_session runner = CliRunner() runner.invoke(goose_cli, ["session", "resume", "session1", "--profile", "default"]) - mock_session_class.assert_called_once_with(name="session1", profile="default") + mock_session_class.assert_called_once_with(name="session1", profile="default", log_level="INFO") mock_session_instance.run.assert_called_once() @@ -59,7 +67,7 @@ def test_session_resume_command_without_session_name_use_latest_session( second_file_path = mock_session_files_path / "second.jsonl" mock_print.assert_called_once_with(f"Resuming most recent session: second from {second_file_path}") - mock_session_class.assert_called_once_with(name="second", profile="default") + mock_session_class.assert_called_once_with(name="second", profile="default", log_level="INFO") mock_session_instance.run.assert_called_once() @@ -121,7 +129,7 @@ def test_combined_group_commands(mock_session): mock_session_class, mock_session_instance = mock_session runner = CliRunner() runner.invoke(cli, ["session", "resume", "session1", "--profile", "default"]) - mock_session_class.assert_called_once_with(name="session1", profile="default") + mock_session_class.assert_called_once_with(name="session1", profile="default", log_level="INFO") mock_session_instance.run.assert_called_once() diff --git a/tests/cli/test_session.py b/tests/cli/test_session.py index 79b6d2bb..6d9086bc 100644 --- a/tests/cli/test_session.py +++ b/tests/cli/test_session.py @@ -1,7 +1,7 @@ from unittest.mock import MagicMock, patch import pytest -from exchange import Message, ToolUse, ToolResult +from exchange import Exchange, Message, ToolUse, ToolResult from goose.cli.prompt.goose_prompt_session import GoosePromptSession from goose.cli.prompt.user_input import PromptAction, UserInput from goose.cli.session import Session @@ -19,12 +19,13 @@ def mock_specified_session_name(): @pytest.fixture def create_session_with_mock_configs(mock_sessions_path, exchange_factory, profile_factory): - with patch("goose.cli.session.build_exchange", return_value=exchange_factory()), patch( + with patch("goose.cli.session.build_exchange") as mock_exchange, patch( "goose.cli.session.load_profile", return_value=profile_factory() ), patch("goose.cli.session.SessionNotifier") as mock_session_notifier, patch( "goose.cli.session.load_provider", return_value="provider" ): mock_session_notifier.return_value = MagicMock() + mock_exchange.return_value = exchange_factory() def create_session(session_attributes: dict = {}): return Session(**session_attributes) @@ -79,59 +80,6 @@ def test_session_removes_tool_use_and_adds_resume_message_if_last_message_is_too ] -def test_save_session_create_session(mock_sessions_path, create_session_with_mock_configs, mock_specified_session_name): - session = create_session_with_mock_configs() - session.exchange.messages.append(Message.assistant("Hello")) - - session.save_session() - session_file = mock_sessions_path / f"{SPECIFIED_SESSION_NAME}.jsonl" - assert session_file.exists() - - saved_messages = session.load_session() - assert len(saved_messages) == 1 - assert saved_messages[0].text == "Hello" - - -def test_save_session_resume_session_new_file( - mock_sessions_path, create_session_with_mock_configs, mock_specified_session_name, create_session_file -): - with patch("goose.cli.session.confirm", return_value=False): - existing_messages = [Message.assistant("existing_message")] - existing_session_file = mock_sessions_path / f"{SESSION_NAME}.jsonl" - create_session_file(existing_messages, existing_session_file) - - new_session_file = mock_sessions_path / f"{SPECIFIED_SESSION_NAME}.jsonl" - assert not new_session_file.exists() - - session = create_session_with_mock_configs({"name": SESSION_NAME}) - session.exchange.messages.append(Message.assistant("new_message")) - - session.save_session() - - assert new_session_file.exists() - assert existing_session_file.exists() - - saved_messages = session.load_session() - assert [message.text for message in saved_messages] == ["existing_message", "new_message"] - - -def test_save_session_resume_session_existing_session_file( - mock_sessions_path, create_session_with_mock_configs, create_session_file -): - with patch("goose.cli.session.confirm", return_value=True): - existing_messages = [Message.assistant("existing_message")] - existing_session_file = mock_sessions_path / f"{SESSION_NAME}.jsonl" - create_session_file(existing_messages, existing_session_file) - - session = create_session_with_mock_configs({"name": SESSION_NAME}) - session.exchange.messages.append(Message.assistant("new_message")) - - session.save_session() - - saved_messages = session.load_session() - assert [message.text for message in saved_messages] == ["existing_message", "new_message"] - - def test_process_first_message_return_message(create_session_with_mock_configs): session = create_session_with_mock_configs() with patch.object( @@ -161,14 +109,6 @@ def test_process_first_message_return_last_exchange_message(create_session_with_ assert len(session.exchange.messages) == 0 -def test_generate_session_name(create_session_with_mock_configs): - session = create_session_with_mock_configs() - with patch.object(GoosePromptSession, "get_save_session_name", return_value=SPECIFIED_SESSION_NAME): - session.generate_session_name() - - assert session.name == SPECIFIED_SESSION_NAME - - def test_log_log_cost(create_session_with_mock_configs): session = create_session_with_mock_configs() mock_logger = MagicMock() @@ -178,3 +118,36 @@ def test_log_log_cost(create_session_with_mock_configs): ), patch("goose.cli.session.get_logger", return_value=mock_logger): session._log_cost() mock_logger.info.assert_called_once_with(cost_message) + + +def test_run_should_auto_save_session(create_session_with_mock_configs, mock_sessions_path): + def custom_exchange_generate(self, *args, **kwargs): + message = Message.assistant("Response") + self.add(message) + return message + + user_inputs = [ + UserInput(action=PromptAction.CONTINUE, text="Question1"), + UserInput(action=PromptAction.CONTINUE, text="Question2"), + UserInput(action=PromptAction.EXIT), + ] + + session = create_session_with_mock_configs({"name": SESSION_NAME}) + with patch.object(GoosePromptSession, "get_user_input", side_effect=user_inputs), patch.object( + Exchange, "generate" + ) as mock_generate, patch("goose.cli.session.save_latest_session") as mock_save_latest_session: + mock_generate.side_effect = lambda *args, **kwargs: custom_exchange_generate(session.exchange, *args, **kwargs) + session.run() + + session_file = mock_sessions_path / f"{SESSION_NAME}.jsonl" + assert session.exchange.generate.call_count == 2 + assert mock_save_latest_session.call_count == 2 + assert mock_save_latest_session.call_args_list[0][0][0] == session_file + assert session_file.exists() + + +def test_set_generated_session_name(create_session_with_mock_configs, mock_sessions_path): + generated_session_name = "generated_session_name" + with patch("goose.cli.session.droid", return_value=generated_session_name): + session = create_session_with_mock_configs({"name": None}) + assert session.name == generated_session_name diff --git a/tests/utils/test_session_file.py b/tests/utils/test_session_file.py index d922bd81..6a2a6498 100644 --- a/tests/utils/test_session_file.py +++ b/tests/utils/test_session_file.py @@ -1,8 +1,16 @@ +import os from pathlib import Path import pytest from exchange import Message -from goose.utils.session_file import list_sorted_session_files, read_from_file, session_file_exists, write_to_file +from goose.utils.session_file import ( + list_sorted_session_files, + read_from_file, + read_or_create_file, + save_latest_session, + session_file_exists, + write_to_file, +) @pytest.fixture @@ -32,6 +40,23 @@ def test_read_from_file_non_jsonl_file(file_path): read_from_file(file_path) +def test_read_or_create_file_when_file_not_exist(tmp_path): + file_path = tmp_path / "no_existing.json" + + assert read_or_create_file(file_path) == [] + assert os.path.exists(file_path) + + +def test_read_or_create_file_when_file_exists(file_path): + messages = [ + Message.user("prompt1"), + ] + write_to_file(file_path, messages) + + assert file_path.exists() + assert read_from_file(file_path) == messages + + def test_list_sorted_session_files(tmp_path): session_files_directory = tmp_path / "session_files_dir" session_files_directory.mkdir() @@ -71,6 +96,21 @@ def test_session_file_exists_return_true_when_session_file_exists(tmp_path): assert session_file_exists(session_files_directory) +def test_save_latest_session(file_path, tmp_path): + messages = [ + Message.user("prompt1"), + Message.user("prompt2"), + ] + write_to_file(file_path, messages) + + messages.append(Message.user("prompt3")) + save_latest_session(file_path, messages) + + messages_in_file = read_from_file(file_path) + assert messages_in_file == messages + assert len(messages_in_file) == 3 + + def create_session_file(file_path, file_name) -> Path: file = file_path / f"{file_name}.jsonl" file.touch()