feat: add guards to session management (#101)
This commit is contained in:
@@ -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
@@ -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__":
|
||||
|
||||
@@ -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)
|
||||
|
||||
Reference in New Issue
Block a user