feat: add guards to session management (#101)

This commit is contained in:
Lam Chau
2024-10-10 05:01:04 -07:00
committed by GitHub
parent 798f346c5a
commit 4375e2fe5e
5 changed files with 184 additions and 34 deletions
@@ -0,0 +1,30 @@
from typing import Any
from rich.prompt import Prompt
class OverwriteSessionPrompt(Prompt):
def __init__(self, *args: tuple[Any], **kwargs: dict[str, Any]) -> None:
super().__init__(*args, **kwargs)
self.choices = {
"yes": "Overwrite the existing session",
"no": "Pick a new session name",
"resume": "Resume the existing session",
}
self.default = "resume"
def check_choice(self, choice: str) -> bool:
for key in self.choices:
normalized_choice = choice.lower()
if normalized_choice == key or normalized_choice[0] == key[0]:
return True
return False
def pre_prompt(self) -> str:
print("Would you like to overwrite it?")
print()
for key, value in self.choices.items():
first_letter, remaining = key[0], key[1:]
rendered_key = rf"[{first_letter}]{remaining}"
print(f" {rendered_key:10} {value}")
print()
+59 -12
View File
@@ -1,22 +1,24 @@
import traceback
from pathlib import Path
from typing import Any, Dict, List, Optional
from typing import Any, Optional
from exchange import Message, ToolResult, ToolUse, Text
from exchange import Message, Text, ToolResult, ToolUse
from rich import print
from rich.markdown import Markdown
from rich.panel import Panel
from rich.prompt import Prompt
from rich.status import Status
from goose.cli.config import ensure_config, session_path, LOG_PATH
from goose._logger import get_logger, setup_logging
from goose.cli.config import LOG_PATH, ensure_config, session_path
from goose.cli.prompt.goose_prompt_session import GoosePromptSession
from goose.cli.session_notifier import SessionNotifier
from goose.cli.prompt.overwrite_session_prompt import OverwriteSessionPrompt
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._create_exchange import create_exchange
from goose.utils.session_file import read_or_create_file, save_latest_session
from goose.utils.session_file import is_empty_session, is_existing_session, read_or_create_file, save_latest_session
RESUME_MESSAGE = "I see we were interrupted. How can I help you?"
@@ -60,7 +62,7 @@ class Session:
profile: Optional[str] = None,
plan: Optional[dict] = None,
log_level: Optional[str] = "INFO",
**kwargs: Dict[str, Any],
**kwargs: dict[str, Any],
) -> None:
if name is None:
self.name = droid()
@@ -69,7 +71,7 @@ class Session:
self.profile_name = profile
self.prompt_session = GoosePromptSession()
self.status_indicator = Status("", spinner="dots")
self.notifier = SessionNotifier(self.status_indicator)
self.notifier = Notifier(self.status_indicator)
self.exchange = create_exchange(profile=load_profile(profile), notifier=self.notifier)
setup_logging(log_file_directory=LOG_PATH, log_level=log_level)
@@ -81,7 +83,7 @@ class Session:
self.prompt_session = GoosePromptSession()
def _get_initial_messages(self) -> List[Message]:
def _get_initial_messages(self) -> list[Message]:
messages = self.load_session()
if messages and messages[-1].role == "user":
@@ -151,8 +153,11 @@ class Session:
Runs the main loop to handle user inputs and responses.
Continues until an empty string is returned from the prompt.
"""
print(f"[dim]starting session | name:[cyan]{self.name}[/] profile:[cyan]{self.profile_name or 'default'}[/]")
print(f"[dim]saving to {self.session_file_path}")
if is_existing_session(self.session_file_path):
self._prompt_overwrite_session()
profile_name = self.profile_name or "default"
print(f"[dim]starting session | name: [cyan]{self.name}[/cyan] profile: [cyan]{profile_name}[/cyan][/dim]")
print()
message = self.process_first_message()
while message: # Loop until no input (empty string).
@@ -178,6 +183,7 @@ class Session:
user_input = self.prompt_session.get_user_input()
message = Message.user(text=user_input.text) if user_input.to_continue() else None
self._remove_empty_session()
self._log_cost()
def reply(self) -> None:
@@ -234,12 +240,53 @@ class Session:
def session_file_path(self) -> Path:
return session_path(self.name)
def load_session(self) -> List[Message]:
def load_session(self) -> list[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()))
print(f"[dim]you can view the cost and token usage in the log directory {LOG_PATH}")
print(f"[dim]you can view the cost and token usage in the log directory {LOG_PATH}[/]")
def _prompt_overwrite_session(self) -> None:
print(f"[yellow]Session already exists at {self.session_file_path}.[/]")
choice = OverwriteSessionPrompt.ask("Enter your choice", show_choices=False)
match choice:
case "y" | "yes":
print("Overwriting existing session")
case "n" | "no":
while True:
new_session_name = Prompt.ask("Enter a new session name")
if not is_existing_session(session_path(new_session_name)):
self.name = new_session_name
break
print(f"[yellow]Session '{new_session_name}' already exists[/]")
case "r" | "resume":
self.exchange.messages.extend(self.load_session())
def _remove_empty_session(self) -> bool:
"""
Removes the session file only when it's empty.
Note: This is because a session file is created at the start of the run
loop. When a user aborts before their first message empty session files
will be created, causing confusion when resuming sessions (which
depends on most recent mtime and is non-empty).
Returns:
bool: True if the session file was removed, False otherwise.
"""
logger = get_logger()
try:
if is_empty_session(self.session_file_path):
logger.debug(f"deleting empty session file: {self.session_file_path}")
self.session_file_path.unlink()
return True
except Exception as e:
logger.error(f"error deleting empty session file: {e}")
return False
if __name__ == "__main__":
+9 -1
View File
@@ -1,7 +1,7 @@
import json
import os
from pathlib import Path
import tempfile
from pathlib import Path
from typing import Dict, Iterator, List
from exchange import Message
@@ -9,6 +9,14 @@ from exchange import Message
from goose.cli.config import SESSION_FILE_SUFFIX
def is_existing_session(path: Path) -> bool:
return path.is_file() and path.stat().st_size > 0
def is_empty_session(path: Path) -> bool:
return path.is_file() and path.stat().st_size == 0
def write_to_file(file_path: Path, messages: List[Message]) -> None:
with open(file_path, "w") as f:
_write_messages_to_file(f, messages)