diff --git a/src/cai/repl/commands/compact.py b/src/cai/repl/commands/compact.py
index af2fc0e9..00df0408 100644
--- a/src/cai/repl/commands/compact.py
+++ b/src/cai/repl/commands/compact.py
@@ -573,10 +573,7 @@ class CompactCommand(Command):
f"\n[#9aa0a6][CAI] Compact current conversation? [/]"
f"[bold white]({msg_count} messages)[/bold white]"
)
- if read_repl_yes_no(
- console,
- "[#9aa0a6][CAI] Compact conversation? [/][bold #00ff9d](y/N): [/]",
- ):
+ if read_repl_yes_no(console, "Compact conversation?"):
# Pass the detected agent name to _perform_compaction
return self._perform_compaction(None, None, agent_name=agent_name)
else:
diff --git a/src/cai/repl/ui/tty_input.py b/src/cai/repl/ui/tty_input.py
index 6e9dc63e..9f801635 100644
--- a/src/cai/repl/ui/tty_input.py
+++ b/src/cai/repl/ui/tty_input.py
@@ -7,13 +7,14 @@ from typing import TYPE_CHECKING
if TYPE_CHECKING:
from rich.console import Console
+_CAI_GREY = "#9aa0a6"
+_CAI_GREEN = "#00ff9d"
-def prepare_tty_for_line_input() -> None:
- """Restore cooked TTY before ``input()`` / Rich ``console.input``."""
- from cai.util.streaming import ensure_cooked_tty, restore_terminal_state
+
+def _restore_tty_after_prompt() -> None:
+ from cai.util.streaming import restore_terminal_state
restore_terminal_state(emit_trailing_newline=False)
- ensure_cooked_tty()
def normalize_repl_line(value: str) -> str:
@@ -27,26 +28,52 @@ def read_repl_line(
*,
markup: bool = True,
) -> str:
- """Read one line with Rich prompt styling; safe after the main CAI> prompt."""
- prepare_tty_for_line_input()
- if prompt:
- console.print(prompt, markup=markup, emoji=False, end="")
+ """Read one line using prompt_toolkit (same stack as the CAI> prompt)."""
+ from prompt_toolkit import prompt as ptk_prompt
+ from prompt_toolkit.formatted_text import HTML
+
+ _restore_tty_after_prompt()
try:
- raw = input()
- except EOFError:
+ if prompt and markup:
+ # Rich markup is for scrollback context only; ptk owns the input line.
+ console.print(prompt, markup=markup, emoji=False, end="")
+ raw = ptk_prompt("")
+ elif prompt:
+ raw = ptk_prompt(prompt)
+ else:
+ raw = ptk_prompt("")
+ except (EOFError, KeyboardInterrupt):
return ""
+ finally:
+ _restore_tty_after_prompt()
return normalize_repl_line(raw)
def read_repl_yes_no(
console: Console,
- prompt: str,
+ label: str,
*,
default: bool = False,
markup: bool = True,
) -> bool:
"""Return True when the user answers y/yes (default answer when they press Enter)."""
- answer = read_repl_line(console, prompt, markup=markup)
+ from prompt_toolkit import prompt as ptk_prompt
+ from prompt_toolkit.formatted_text import HTML
+
+ _restore_tty_after_prompt()
+ suffix = "Y/n" if default else "y/N"
+ ptk_prompt_text = HTML(
+ f' {label} '
+ f': '
+ )
+ try:
+ raw = ptk_prompt(ptk_prompt_text)
+ except (EOFError, KeyboardInterrupt):
+ return default
+ finally:
+ _restore_tty_after_prompt()
+
+ answer = normalize_repl_line(raw)
if not answer:
return default
return answer.lower() in ("y", "yes")
diff --git a/tests/repl/test_tty_input.py b/tests/repl/test_tty_input.py
index 13af3e2c..0d44ec57 100644
--- a/tests/repl/test_tty_input.py
+++ b/tests/repl/test_tty_input.py
@@ -20,29 +20,29 @@ class TestNormalizeReplLine:
class TestReadReplYesNo:
- @patch("cai.repl.ui.tty_input.input", return_value="y\r")
- @patch("cai.repl.ui.tty_input.prepare_tty_for_line_input")
- def test_yes_with_carriage_return(self, _prepare, _input):
+ @patch("prompt_toolkit.prompt", return_value="y\r")
+ @patch("cai.repl.ui.tty_input._restore_tty_after_prompt")
+ def test_yes_with_carriage_return(self, _restore, _prompt):
console = MagicMock()
- assert read_repl_yes_no(console, "Continue? (y/N): ") is True
+ assert read_repl_yes_no(console, "Continue") is True
- @patch("cai.repl.ui.tty_input.input", return_value="")
- @patch("cai.repl.ui.tty_input.prepare_tty_for_line_input")
- def test_empty_defaults_to_no(self, _prepare, _input):
+ @patch("prompt_toolkit.prompt", return_value="")
+ @patch("cai.repl.ui.tty_input._restore_tty_after_prompt")
+ def test_empty_defaults_to_no(self, _restore, _prompt):
console = MagicMock()
- assert read_repl_yes_no(console, "Continue? (y/N): ", default=False) is False
+ assert read_repl_yes_no(console, "Continue", default=False) is False
- @patch("cai.repl.ui.tty_input.input", return_value="n")
- @patch("cai.repl.ui.tty_input.prepare_tty_for_line_input")
- def test_no_answer(self, _prepare, _input):
+ @patch("prompt_toolkit.prompt", return_value="n")
+ @patch("cai.repl.ui.tty_input._restore_tty_after_prompt")
+ def test_no_answer(self, _restore, _prompt):
console = MagicMock()
- assert read_repl_yes_no(console, "Continue? (y/N): ") is False
+ assert read_repl_yes_no(console, "Continue") is False
class TestReadReplLine:
- @patch("cai.repl.ui.tty_input.input", return_value="RESET\r")
- @patch("cai.repl.ui.tty_input.prepare_tty_for_line_input")
- def test_prepares_tty_before_read(self, prepare, _input):
+ @patch("prompt_toolkit.prompt", return_value="RESET\r")
+ @patch("cai.repl.ui.tty_input._restore_tty_after_prompt")
+ def test_restores_tty_around_prompt(self, restore, _prompt):
console = MagicMock()
- assert read_repl_line(console, "> ") == "RESET"
- prepare.assert_called_once()
+ assert read_repl_line(console, "> ", markup=False) == "RESET"
+ assert restore.call_count == 2