Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
34 changes: 28 additions & 6 deletions src/tau_coding/tui/app.py
Original file line number Diff line number Diff line change
Expand Up @@ -4766,8 +4766,27 @@ def on_text_area_changed(self, event: TextArea.Changed) -> None:
return
prompt = self.query_one("#prompt", PromptInput)
prompt.sync_pending_paste()
self._sync_prompt_shell_mode(event.text_area.text)
self._completion_state = self._build_completion_state(event.text_area.text)
# Read text and cursor from the widget so both come from one snapshot.
text = prompt.text
self._sync_prompt_shell_mode(text)
self._completion_state = self._build_completion_state(text, cursor=prompt.cursor_position)
self._refresh_completions()

def on_text_area_selection_changed(self, event: TextArea.SelectionChanged) -> None:
"""Close prompt autocomplete when the caret leaves the completed token."""
if event.text_area.id != "prompt":
return
# Edits post SelectionChanged before Changed; check after Changed has rebuilt.
self.call_later(self._close_completions_if_caret_left_token)

def _close_completions_if_caret_left_token(self) -> None:
if not self._completion_state.items:
return
item = self._completion_state.items[0]
cursor = self.query_one("#prompt", PromptInput).cursor_position
if item.start < cursor <= item.end:
return
self._completion_state = CompletionState()
self._refresh_completions()

async def action_submit_prompt(self) -> None:
Expand Down Expand Up @@ -6193,12 +6212,14 @@ def action_accept_completion(self) -> None:
self.screen.action_select_cursor()
return
prompt = self.query_one("#prompt", PromptInput)
item = self._completion_state.selected
applied = self._apply_selected_completion(prompt.text)
if applied is None:
if applied is None or item is None:
return
prompt.text = applied
prompt.move_cursor(_text_end_location(applied))
self._completion_state = self._build_completion_state(prompt.text)
cursor = item.cursor_after_apply()
prompt.cursor_position = cursor
self._completion_state = self._build_completion_state(prompt.text, cursor=cursor)
self._refresh_completions()

def action_completion_next(self) -> None:
Expand Down Expand Up @@ -7431,10 +7452,11 @@ def _toggle_sidebar_visibility(self) -> None:
state = "shown" if self._sidebar_visibility_override else "hidden"
self._notify(f"Sidebar {state} for this session.")

def _build_completion_state(self, text: str) -> CompletionState:
def _build_completion_state(self, text: str, *, cursor: int | None = None) -> CompletionState:
registry = _session_command_registry(self.session)
return build_completion_state(
text,
cursor=cursor,
command_registry=registry,
skills=self.session.skills,
prompt_templates=self.session.prompt_templates,
Expand Down
70 changes: 51 additions & 19 deletions src/tau_coding/tui/autocomplete.py
Original file line number Diff line number Diff line change
Expand Up @@ -65,6 +65,10 @@ def apply(self, text: str) -> str:
"""Apply this completion to input text."""
return f"{text[: self.start]}{self.replacement}{text[self.end :]}"

def cursor_after_apply(self) -> int:
"""Return the cursor offset just after the applied replacement."""
return self.start + len(self.replacement)


@dataclass(frozen=True, slots=True)
class CompletionState:
Expand Down Expand Up @@ -102,6 +106,7 @@ def select_previous(self) -> CompletionState:
def build_completion_state(
text: str,
*,
cursor: int | None = None,
command_registry: CommandRegistry,
skills: Sequence[Skill],
prompt_templates: Sequence[PromptTemplate],
Expand All @@ -114,12 +119,13 @@ def build_completion_state(
cwd: Path | None = None,
) -> CompletionState:
"""Build autocomplete suggestions for the current prompt text."""
cursor = len(text) if cursor is None else max(0, min(cursor, len(text)))
if not text.startswith("/") or text.startswith("//"):
if cwd is not None:
shell_completions = _shell_path_completions(text=text, cwd=cwd)
shell_completions = _shell_path_completions(text=text, cursor=cursor, cwd=cwd)
if shell_completions is not None:
return CompletionState(shell_completions)
return CompletionState(_file_reference_completions(text=text, cwd=cwd))
return CompletionState(_file_reference_completions(text=text, cursor=cursor, cwd=cwd))
return CompletionState()

token_end = _first_token_end(text)
Expand All @@ -129,7 +135,9 @@ def build_completion_state(
if has_argument_text and _matches_skill_command(token, skills):
# Skill arguments are prompt text, so @ file references stay available.
if cwd is not None:
return CompletionState(_file_reference_completions(text=text, cwd=cwd))
return CompletionState(
_file_reference_completions(text=text, cursor=cursor, cwd=cwd)
)
return CompletionState()
return CompletionState(_skill_completions(token=token, token_end=token_end, skills=skills))

Expand All @@ -151,7 +159,7 @@ def build_completion_state(

if has_argument_text and _matches_prompt_template_command(token, prompt_templates):
if cwd is not None:
return CompletionState(_file_reference_completions(text=text, cwd=cwd))
return CompletionState(_file_reference_completions(text=text, cursor=cursor, cwd=cwd))
return CompletionState()

if has_argument_text and _matches_registered_command(token, command_registry):
Expand All @@ -167,14 +175,16 @@ def build_completion_state(
)


def _file_reference_completions(*, text: str, cwd: Path) -> tuple[CompletionItem, ...]:
token = _active_file_reference_token(text)
def _file_reference_completions(*, text: str, cursor: int, cwd: Path) -> tuple[CompletionItem, ...]:
token = _active_file_reference_token(text, cursor)
if token is None:
return ()
start, end = token
prefix = text[start + 1 : end]
prefix = text[start + 1 : cursor]
existing = text[start:end]
external_completions = _external_file_reference_completions(
prefix=prefix,
existing=existing,
start=start,
end=end,
cwd=cwd,
Expand All @@ -188,6 +198,8 @@ def _file_reference_completions(*, text: str, cwd: Path) -> tuple[CompletionItem
if prefix.lower() not in relative.lower():
continue
display = f"@{relative}{'/' if path.is_dir() else ''}"
if display == existing:
continue
suggestions.append(
CompletionItem(
display=display,
Expand All @@ -204,13 +216,15 @@ def _file_reference_completions(*, text: str, cwd: Path) -> tuple[CompletionItem


def _external_file_reference_completions(
*, prefix: str, start: int, end: int, cwd: Path
*, prefix: str, existing: str, start: int, end: int, cwd: Path
) -> tuple[CompletionItem, ...] | None:
if prefix == ".." or prefix.endswith("/.."):
target = cwd / prefix
if not target.is_dir():
return ()
display = f"@{prefix}/"
if display == existing:
return ()
return (
CompletionItem(
display=display,
Expand Down Expand Up @@ -244,7 +258,7 @@ def _external_file_reference_completions(
if not child.name.lower().startswith(name_prefix.lower()):
continue
display = f"@{parent_text}/{child.name}{'/' if child.is_dir() else ''}"
if display == f"@{prefix}":
if display == existing:
continue
suggestions.append(
CompletionItem(
Expand All @@ -261,13 +275,15 @@ def _external_file_reference_completions(
return tuple(suggestions)


def _active_file_reference_token(text: str) -> tuple[int, int] | None:
cursor = len(text)
def _active_file_reference_token(text: str, cursor: int) -> tuple[int, int] | None:
token_start = max(text.rfind(" ", 0, cursor), text.rfind("\n", 0, cursor)) + 1
at_index = text.rfind("@", token_start, cursor)
if at_index == -1:
return None
return at_index, cursor
end = cursor
while end < len(text) and not text[end].isspace():
end += 1
return at_index, end


def _iter_file_reference_paths(cwd: Path) -> tuple[Path, ...]:
Expand Down Expand Up @@ -298,13 +314,17 @@ def _is_ignored_file_completion_path(path: Path, *, cwd: Path) -> bool:
return any(part in IGNORED_FILE_COMPLETION_DIRS for part in relative_parts)


def _shell_path_completions(*, text: str, cwd: Path) -> tuple[CompletionItem, ...] | None:
def _shell_path_completions(
*, text: str, cursor: int, cwd: Path
) -> tuple[CompletionItem, ...] | None:
prefix_span = _shell_command_prefix_span(text)
if prefix_span is None:
return None
if cursor < prefix_span[1]:
return ()

start, end = _active_shell_path_token(text=text, command_start=prefix_span[1])
token = text[start:end]
start, end = _active_shell_path_token(text=text, cursor=cursor, command_start=prefix_span[1])
token = text[start:cursor]
if not token:
return ()

Expand Down Expand Up @@ -332,7 +352,7 @@ def _shell_path_completions(*, text: str, cwd: Path) -> tuple[CompletionItem, ..
continue
relative = child.relative_to(cwd).as_posix()
replacement = f"{replacement_prefix}{relative}{'/' if child.is_dir() else ''}"
if replacement == token:
if replacement == text[start:end]:
continue
suggestions.append(
CompletionItem(
Expand All @@ -359,8 +379,7 @@ def _shell_command_prefix_span(text: str) -> tuple[int, int] | None:
return None


def _active_shell_path_token(*, text: str, command_start: int) -> tuple[int, int]:
cursor = len(text)
def _active_shell_path_token(*, text: str, cursor: int, command_start: int) -> tuple[int, int]:
token_start = command_start
escaped = False
for index in range(cursor - 1, command_start - 1, -1):
Expand All @@ -374,7 +393,20 @@ def _active_shell_path_token(*, text: str, command_start: int) -> tuple[int, int
if char.isspace():
token_start = index + 1
break
return token_start, cursor
return token_start, _shell_token_end(text, cursor)


def _shell_token_end(text: str, cursor: int) -> int:
index = cursor
while index < len(text):
char = text[index]
if char == "\\":
index += 2
continue
if char.isspace():
return index
index += 1
return len(text)


def _parse_shell_path_token(token: str) -> tuple[str, str, str] | None:
Expand Down
78 changes: 78 additions & 0 deletions tests/test_tui_app.py
Original file line number Diff line number Diff line change
Expand Up @@ -3056,6 +3056,84 @@ async def test_tui_submit_multiple_large_pastes_sends_all_full_content() -> None
assert session.prompt_texts == [f"{first}\nthen\n{second}"]


@pytest.mark.anyio
async def test_tui_caret_leaving_mention_closes_completions_but_does_not_open(
tmp_path: Path,
) -> None:
(tmp_path / "src").mkdir()
(tmp_path / "src" / "app.py").write_text("print('hi')\n", encoding="utf-8")
session = FakeSession()
session.cwd = tmp_path
app = TauTuiApp(session)

async with app.run_test() as pilot:
prompt = app.query_one("#prompt", PromptInput)
prompt.text = "look at @a and fix it"
prompt.cursor_position = len("look at @a")
prompt.insert("p")
await pilot.pause()
assert [item.display for item in app._completion_state.items] == ["@src/app.py"]

prompt.cursor_position = len("look at @a")
await pilot.pause()
assert [item.display for item in app._completion_state.items] == ["@src/app.py"]

prompt.cursor_position = len("look at @ap and")
await pilot.pause()
assert app._completion_state.items == ()

prompt.cursor_position = len("look at @ap")
await pilot.pause()
assert app._completion_state.items == ()


@pytest.mark.anyio
async def test_tui_enter_submits_with_caret_parked_in_complete_mention(
tmp_path: Path,
) -> None:
(tmp_path / "src").mkdir()
(tmp_path / "src" / "app.py").write_text("print('hi')\n", encoding="utf-8")
session = FakeSession()
session.cwd = tmp_path
app = TauTuiApp(session)

async with app.run_test() as pilot:
prompt = app.query_one("#prompt", PromptInput)
prompt.text = "look at @src/app.py and fix"
await pilot.pause()
prompt.cursor_position = len("look at @src")
await pilot.pause()
assert app._completion_state.items == ()
await pilot.press("enter")
await pilot.pause()

assert session.prompt_texts == ["look at @src/app.py and fix"]


@pytest.mark.anyio
async def test_tui_accepting_mid_prompt_completion_keeps_cursor_after_mention(
tmp_path: Path,
) -> None:
(tmp_path / "src").mkdir()
(tmp_path / "src" / "app.py").write_text("print('hi')\n", encoding="utf-8")
session = FakeSession()
session.cwd = tmp_path
app = TauTuiApp(session)

async with app.run_test() as pilot:
prompt = app.query_one("#prompt", PromptInput)
prompt.text = "look at @ap and fix it"
prompt.cursor_position = len("look at @ap")
prompt.insert("p")
await pilot.pause()
app.action_accept_completion()
await pilot.pause()

assert prompt.text == "look at @src/app.py and fix it"
assert prompt.cursor_position == len("look at @src/app.py")
assert app._completion_state.items == ()


@pytest.mark.anyio
async def test_tui_app_mounts_sidebar_and_transcript() -> None:
app = TauTuiApp(FakeSession())
Expand Down
Loading