diff --git a/src/lecode/agent/runner.py b/src/lecode/agent/runner.py index 208da4c..82203f6 100644 --- a/src/lecode/agent/runner.py +++ b/src/lecode/agent/runner.py @@ -304,6 +304,7 @@ async def run( replay = self.store.load_for_model(self.session) if replay and replay[0].get("role") == "system" and history[:1] == replay[:1]: history.insert(0, {"role": "system", "content": ""}) + del replay # The live conversation, visible through ctx (subagents, hooks). self.ctx.extras["conversation"] = history manager = self.ctx.extras.get("workers") diff --git a/src/lecode/agent/tools/bash.py b/src/lecode/agent/tools/bash.py index e265bc8..31bd319 100644 --- a/src/lecode/agent/tools/bash.py +++ b/src/lecode/agent/tools/bash.py @@ -11,12 +11,14 @@ from __future__ import annotations import asyncio +import codecs import contextlib +import tempfile import time import uuid from collections.abc import Callable from pathlib import Path -from typing import Any +from typing import Any, BinaryIO from lecode.agent.tools.base import Tool, ToolContext, ToolResult from lecode.extras.proc import ProcResult @@ -71,7 +73,7 @@ async def _run_shell( buffer = bytearray() deadline = time.monotonic() + timeout idle_deadline = time.monotonic() + idle_timeout - timed_out = idle_killed = False + timed_out = idle_killed = aborted = False try: while True: @@ -79,14 +81,12 @@ async def _run_shell( if wait <= 0: timed_out = time.monotonic() >= deadline idle_killed = not timed_out - _kill_tree(proc) break try: chunk = await asyncio.wait_for(proc.stdout.read(65536), timeout=wait) except TimeoutError: timed_out = time.monotonic() >= deadline idle_killed = not timed_out - _kill_tree(proc) break if not chunk: # EOF: process exited and pipes drained break @@ -96,12 +96,22 @@ async def _run_shell( if len(buffer) > max_bytes: del buffer[: len(buffer) - max_bytes] idle_deadline = time.monotonic() + idle_timeout - except asyncio.CancelledError: - # Turn aborted (Ctrl-C): never leave the child running. - _kill_tree(proc) - with contextlib.suppress(TimeoutError, asyncio.CancelledError): - await asyncio.wait_for(proc.wait(), timeout=5.0) + except BaseException: + aborted = True raise + finally: + if aborted or timed_out or idle_killed: + _kill_tree(proc) + with contextlib.suppress(TimeoutError, asyncio.CancelledError): + async with asyncio.timeout(5.0): + # wait() can wait on pipes held by children missed during a fork. + while proc.returncode is None: + await asyncio.sleep(0.01) + # Catch children spawned while the first signal killed the shell. + _kill_tree(proc) + while await proc.stdout.read(65536): + pass + await proc.wait() with contextlib.suppress(TimeoutError): await asyncio.wait_for(proc.wait(), timeout=5.0) @@ -109,14 +119,21 @@ async def _run_shell( return bytes(buffer), exit_code, timed_out, idle_killed -def _save_overflow(text: str) -> Path: +def _save_overflow(output: BinaryIO) -> tuple[Path, str, str]: from lecode.config.loader import config_dir overflow_dir = config_dir() / "overflow" overflow_dir.mkdir(parents=True, exist_ok=True) path = overflow_dir / f"{uuid.uuid4().hex}.log" - path.write_text(text, encoding="utf-8") - return path + head = tail = "" + half = MAX_OUTPUT_BYTES // 2 + chunks = iter(lambda: output.read(65536), b"") + with path.open("w", encoding="utf-8") as destination: + for text in codecs.iterdecode(chunks, "utf-8", errors="replace"): + destination.write(text) + head += text[: half - len(head)] + tail = (tail + text)[-half:] + return path, head, tail class BashTool(Tool): @@ -185,22 +202,24 @@ async def run(self, args: dict, ctx: ToolContext) -> ToolResult: return ToolResult(f"error: {e}", is_error=True) return ToolResult(f"background task {record.id} started: {command}") try: - output, exit_code, timed_out, idle_killed = await _run_shell( - command, ctx.cwd, timeout, idle, MAX_OUTPUT_BYTES - ) + with tempfile.SpooledTemporaryFile(max_size=MAX_OUTPUT_BYTES) as output: + _, exit_code, timed_out, idle_killed = await _run_shell( + command, ctx.cwd, timeout, idle, MAX_OUTPUT_BYTES, on_chunk=output.write + ) + truncated = output.tell() > MAX_OUTPUT_BYTES + output.seek(0) + if truncated: + full_path, head, tail = _save_overflow(output) + text = ( + head + + f"\n… [output truncated; full output saved to {full_path}] …\n" + + tail + ) + else: + text = output.read().decode("utf-8", errors="replace") except OSError as e: return ToolResult(f"error: {e}", is_error=True) - text = output.decode("utf-8", errors="replace") - if len(output) > MAX_OUTPUT_BYTES: - full_path = _save_overflow(text) - half = MAX_OUTPUT_BYTES // 2 - text = ( - text[:half] - + f"\n… [output truncated; full output saved to {full_path}] …\n" - + text[-half:] - ) - notes: list[str] = [] if timed_out: notes.append(f"timed out after {timeout}s") @@ -218,7 +237,7 @@ async def run(self, args: dict, ctx: ToolContext) -> ToolResult: stdout=text, stderr="", # The shell merges stderr into stdout. timed_out=timed_out or idle_killed, - truncated=len(output) > MAX_OUTPUT_BYTES, + truncated=truncated, ) }, ) diff --git a/src/lecode/agent/tools/read.py b/src/lecode/agent/tools/read.py index 330b348..0a586a2 100644 --- a/src/lecode/agent/tools/read.py +++ b/src/lecode/agent/tools/read.py @@ -9,6 +9,7 @@ from __future__ import annotations import zlib +from itertools import islice from pathlib import Path from lecode.agent.tools.base import Tool, ToolContext, ToolResult @@ -97,12 +98,22 @@ async def run(self, args: dict, ctx: ToolContext) -> ToolResult: with_anchors = bool(args.get("with_anchors")) try: - lines = path.read_text(encoding="utf-8", errors="replace").splitlines() + with path.open(encoding="utf-8", errors="replace") as stream: + lines = (text for line in stream for text in line.splitlines()) + page: list[str] = [] + total = 0 + for total, text in enumerate(lines, start=1): + if offset <= total < offset + limit: + page.append(text) + if limit < 0: + # Preserve negative slice bounds without retaining the whole file. + start, stop, _ = slice(offset - 1, offset - 1 + limit).indices(total) + stream.seek(0) + lines = (text for line in stream for text in line.splitlines()) + page = list(islice(lines, start, stop)) except OSError as e: return ToolResult(f"error: {e}", is_error=True) - total = len(lines) - page = lines[offset - 1 : offset - 1 + limit] with_marks: list[str] = [] for i, text in enumerate(page, start=offset): if with_anchors: diff --git a/src/lecode/cli.py b/src/lecode/cli.py index a9aa1ae..b6c4f4a 100644 --- a/src/lecode/cli.py +++ b/src/lecode/cli.py @@ -14,7 +14,7 @@ import sys from collections.abc import Callable, Coroutine from pathlib import Path -from typing import Annotated, Any +from typing import TYPE_CHECKING, Annotated, Any import typer from typer.core import TyperGroup, TyperOption @@ -63,10 +63,10 @@ SessionNotFoundError, SessionStore, ) -from lecode.setup_wizard import offer_first_run_setup, run_wizard from lecode.telemetry import init_telemetry, shutdown_telemetry -from lecode.tui.app import TuiApp -from lecode.tui.name_prompt import prompt_session_name + +if TYPE_CHECKING: + from lecode.tui.app import TuiApp #: Exit codes (headless mode uses the same taxonomy). EXIT_OK = 0 @@ -688,6 +688,8 @@ def run_interactive( """ from rich.console import Console + from lecode.setup_wizard import offer_first_run_setup + from lecode.tui.app import TuiApp from lecode.tui.loading import ( LoadingProgress, build_load_report, @@ -695,6 +697,7 @@ def run_interactive( provider_step, session_step, ) + from lecode.tui.name_prompt import prompt_session_name # Banner first — before any slow work (setup wizard, network fetches). console = Console(no_color=no_color) @@ -915,6 +918,8 @@ def run_setup() -> int: if not sys.stdin.isatty(): typer.echo("error: --setup requires an interactive terminal", err=True) return EXIT_STARTUP + from lecode.setup_wizard import run_wizard + try: path = asyncio.run(run_wizard()) except (KeyboardInterrupt, EOFError): diff --git a/src/lecode/extras/proc.py b/src/lecode/extras/proc.py index 6ca1971..5af5462 100644 --- a/src/lecode/extras/proc.py +++ b/src/lecode/extras/proc.py @@ -36,20 +36,27 @@ class ProcResult: truncated: bool = False -def _cap(data: bytes, max_bytes: int) -> tuple[str, bool]: - """Decode with head/tail keeping; the middle is elided past the cap.""" - if len(data) <= max_bytes: - return data.decode("utf-8", errors="replace"), False +async def _read_capped(stream: asyncio.StreamReader, max_bytes: int) -> tuple[str, bool]: + """Drain a pipe while retaining the reported prefix and suffix.""" + head = bytearray() + tail = bytearray() + total = 0 half = max_bytes // 2 - head = data[:half] - tail = data[-half:] - skipped = len(data) - len(head) - len(tail) - text = ( - head.decode("utf-8", errors="replace") - + TRUNCATION_MARKER.format(skipped=skipped) - + tail.decode("utf-8", errors="replace") + while chunk := await stream.read(65536): + total += len(chunk) + head.extend(chunk[: max_bytes - len(head)]) + if half: + tail.extend(chunk) + if len(tail) > half: + del tail[:-half] + if total <= max_bytes: + return bytes(head).decode("utf-8", errors="replace"), False + return ( + bytes(head[:half]).decode("utf-8", errors="replace") + + TRUNCATION_MARKER.format(skipped=total - half - len(tail)) + + bytes(tail).decode("utf-8", errors="replace"), + True, ) - return text, True async def run_proc( @@ -70,27 +77,50 @@ async def run_proc( limit=max_output_bytes * 2, start_new_session=True, # own process group so _kill works on trees ) + + async def feed_input() -> None: + if proc.stdin is not None: + try: + proc.stdin.write(input.encode()) + await proc.stdin.drain() + except (BrokenPipeError, ConnectionResetError): + pass # Match communicate(): a child may close stdin early. + finally: + proc.stdin.close() + + tasks = [ + asyncio.create_task(_read_capped(proc.stdout, max_output_bytes)), + asyncio.create_task(_read_capped(proc.stderr, max_output_bytes)), + asyncio.create_task(feed_input()), + ] + + async def communicate(): + stdout, stderr, _ = await asyncio.gather(*tasks) + await proc.wait() + return stdout, stderr + + communication = asyncio.create_task(communicate()) timed_out = False try: - stdout_b, stderr_b = await asyncio.wait_for( - proc.communicate(input.encode() if input is not None else None), - timeout=timeout, - ) + results = await asyncio.wait_for(asyncio.shield(communication), timeout=timeout) except TimeoutError: timed_out = True _kill(proc) try: - stdout_b, stderr_b = await asyncio.wait_for(proc.communicate(), timeout=5.0) + results = await asyncio.wait_for(asyncio.shield(communication), timeout=5.0) except TimeoutError: - stdout_b, stderr_b = b"", b"" - except asyncio.CancelledError: - # Caller aborted (Ctrl-C): never leave the child running. + results = [("", False), ("", False)] + except BaseException: _kill(proc) with contextlib.suppress(TimeoutError, asyncio.CancelledError): await asyncio.wait_for(proc.wait(), timeout=5.0) raise - stdout, out_truncated = _cap(stdout_b or b"", max_output_bytes) - stderr, err_truncated = _cap(stderr_b or b"", max_output_bytes) + finally: + communication.cancel() + for task in tasks: + task.cancel() + await asyncio.gather(communication, *tasks, return_exceptions=True) + (stdout, out_truncated), (stderr, err_truncated) = results[:2] exit_code = proc.returncode if proc.returncode is not None else -1 return ProcResult( exit_code=exit_code, diff --git a/src/lecode/session/storage.py b/src/lecode/session/storage.py index 9cd0901..fc6297c 100644 --- a/src/lecode/session/storage.py +++ b/src/lecode/session/storage.py @@ -24,7 +24,7 @@ import re import sqlite3 import uuid -from collections.abc import Callable +from collections.abc import Callable, Iterator from dataclasses import asdict, dataclass from datetime import UTC, datetime from pathlib import Path @@ -273,14 +273,22 @@ def open(self, session_id: str) -> Session: raise ValueError("session symlinks are not allowed") if not path.is_file(): raise SessionNotFoundError(session_id) - records = self._read_records_at(path) - meta = next((r for r in records if isinstance(r, MetaRecord)), None) + meta: MetaRecord | None = None + new_name = None + max_seq: int | None = None + for record in self._iter_records_at(path): + if isinstance(record, MetaRecord): + if meta is None: + meta = record + else: + max_seq = record.seq if max_seq is None else max(max_seq, record.seq) + if isinstance(record, EventRecord) and record.kind == "rename": + new_name = record.data.get("name") if meta is None: raise SessionNotFoundError(f"{session_id} (no meta record)") - renames = [r for r in records if isinstance(r, EventRecord) and r.kind == "rename"] - if renames and (new_name := renames[-1].data.get("name")): + if new_name: meta = meta.model_copy(update={"name": str(new_name)}) - return Session(meta=meta, path=path, next_seq=_next_seq(records)) + return Session(meta=meta, path=path, next_seq=1 + (max_seq if max_seq is not None else 0)) def acquire_lock(self, session: Session) -> SessionLock | None: """Lock the session against a second live lecode process. @@ -466,25 +474,34 @@ def source_snapshot( raise ValueError("session symlinks are not allowed") if not path.is_file(): return SourceSnapshot("missing") + records: list[Record] = [] + exact: list[dict[str, Any]] = [] try: - raw = [ - json.loads(line) - for line in path.read_text(encoding="utf-8").splitlines() - if line.strip() - ] - parsed = [parse_record(json.dumps(record)) for record in raw] + with path.open(encoding="utf-8") as source: + lines = (line for chunk in source for line in chunk.splitlines()) + for line in lines: + if not line.strip(): + continue + raw = json.loads(line) + record = parse_record(line) + if record is None or (records and type(raw.get("seq")) is not int): + return SourceSnapshot("stale") + if isinstance(record, MessageRecord): + if start_seq <= record.seq <= end_seq: + exact.append(raw) + else: + # Visibility needs metadata, not unrelated message payloads. + record.message = {} + records.append(record) except (ValueError, UnicodeError, OSError): return SourceSnapshot("stale") - if not parsed or any(record is None for record in parsed): + if not records: return SourceSnapshot("stale") - records = [record for record in parsed if record is not None] meta = records[0] if not isinstance(meta, MetaRecord) or meta.id != session_id: return SourceSnapshot("stale") if any(isinstance(record, MetaRecord) for record in records[1:]): return SourceSnapshot("stale") - if any(type(record.get("seq")) is not int for record in raw[1:]): - return SourceSnapshot("stale") if resolve_project_root(meta.cwd) != resolve_project_root(project_root): raise ValueError("source belongs to another project") excluded = excluded | self.exclusion_reader(session_id) @@ -504,7 +521,6 @@ def source_snapshot( } if any(r.seq not in visible or r.seq in excluded for r in selected): return SourceSnapshot("hidden") - exact = [r for r in raw if r.get("type") == "message" and start_seq <= r["seq"] <= end_seq] digest = hashlib.sha256( json.dumps(exact, sort_keys=True, separators=(",", ":"), ensure_ascii=True).encode() ).hexdigest() @@ -543,8 +559,7 @@ def validate_source( return SourceSnapshot("stale") return snapshot - def _read_records_at(self, path: Path) -> list[Record]: - records: list[Record] = [] + def _iter_records_at(self, path: Path) -> Iterator[Record]: with path.open(encoding="utf-8") as f: for line in f: if not line.strip(): @@ -553,8 +568,10 @@ def _read_records_at(self, path: Path) -> list[Record]: if record is None: self.corrupt_lines += 1 else: - records.append(record) - return records + yield record + + def _read_records_at(self, path: Path) -> list[Record]: + return list(self._iter_records_at(path)) def read_records(self, session: Session) -> list[Record]: """All records in file order; corrupt lines are skipped and counted.""" @@ -1038,6 +1055,6 @@ def load_grants(self, session: Session) -> list[tuple[str, str]]: """All persisted (tool, pattern) permission grants.""" return [ (str(r.data.get("tool", "")), str(r.data.get("pattern", ""))) - for r in self.read_records(session) + for r in self._iter_records_at(session.path) if isinstance(r, EventRecord) and r.kind == "permission_grant" ] diff --git a/tests/test_agent_runner.py b/tests/test_agent_runner.py index f1d0d43..a69cc0a 100644 --- a/tests/test_agent_runner.py +++ b/tests/test_agent_runner.py @@ -4,6 +4,7 @@ import asyncio import json +import weakref from types import SimpleNamespace import pytest @@ -108,6 +109,46 @@ async def test_run_refreshes_system_prompt_in_place(tool_ctx): assert messages[0]["content"] == "REFRESHED PROMPT" +async def test_startup_replay_is_released_before_model_request(tool_ctx, tmp_path, monkeypatch): + store = SessionStore(tmp_path / "cfg") + session = store.create("resume", tool_ctx.cwd) + store.append_message(session, {"role": "user", "content": "prior request"}) + history = store.load_for_model(session) + load = store.load_for_model + replay_ref = None + + class Replay(list): + pass + + def tracked_replay(session): + nonlocal replay_ref + replay = Replay(load(session)) + if replay_ref is None: + replay_ref = weakref.ref(replay) + return replay + + monkeypatch.setattr(store, "load_for_model", tracked_replay) + runner, provider = make_runner( + tool_ctx, + [{"text": "done"}], + session=session, + store=store, + refresh_prompt=lambda: "base instructions", + ) + + async def check_released(event): + if isinstance(event, LlmCall): + assert replay_ref is not None and replay_ref() is None + + result = await runner.run(history, check_released) + assert result.final_text == "done" + assert provider.requests[0]["messages"][:2] == [ + {"role": "system", "content": "base instructions"}, + *history, + ] + store.close() + + async def test_memory_write_refreshes_next_request_preserving_turn_overlay(tool_ctx): state = {"memory": "old memory"} diff --git a/tests/test_headless.py b/tests/test_headless.py index 71fda2d..4cc8ba1 100644 --- a/tests/test_headless.py +++ b/tests/test_headless.py @@ -4,6 +4,9 @@ import json import re +import subprocess +import sys +from pathlib import Path import pytest from tests.fakes import FakeProvider @@ -112,3 +115,63 @@ def test_no_args_launches_interactive(headless, monkeypatch): result = runner.invoke(app, []) assert result.exit_code == EXIT_OK assert len(calls) == 1 + + +@pytest.mark.parametrize("case", ["short", "tools"]) +def test_headless_does_not_load_interactive_modules(tmp_path, case): + root = Path(__file__).resolve().parents[1] + script = """ +import os +import sys +from pathlib import Path +from types import SimpleNamespace + +from lecode import cli +from lecode.config.models import Config +from tests.fakes import FakeProvider + +project = Path(sys.argv[1]) / "project" +project.mkdir() +(project / ".git").mkdir() +(project / "input.txt").write_text("line\\n" * 64) +os.chdir(project) +os.environ["LECODE_CONFIG_DIR"] = str(project / "cfg") +os.environ["LECODE_SKILLS_DIR"] = str(project / "skills") +for name in ("HERDR_ENV", "HERDR_BIN_PATH", "HERDR_PANE_ID"): + os.environ.pop(name, None) +config = Config() +config.mcp.enable_exa = False +config.mcp.enable_context7 = False +cli.load_config = lambda: SimpleNamespace(config=config) +turns = 20 if sys.argv[2] == "tools" else 0 +provider = FakeProvider([ + {"tool_calls": [{"id": f"read-{i}", "name": "read", + "arguments": f'{{"path":"input.txt","offset":{i},"limit":10}}'}]} + for i in range(1, turns + 1) +] + [{"text": "final answer"}]) +cli.build_provider = lambda config, api_key=None: provider +cli.check_dependencies = lambda: None +cli.app(args=["-p", "hello"], standalone_mode=False) +assert len(provider.requests) == turns + 1 +results = [m for m in provider.requests[-1]["messages"] if m["role"] == "tool"] +assert len(results) == turns +if results: + assert results[-1]["content"].startswith("20\\tline") +assert not any(m.startswith("lecode.tui.") for m in sys.modules) +assert "prompt_toolkit" not in sys.modules +""" + result = subprocess.run( + [ + sys.executable, + "-c", + script, + str(tmp_path), + case, + ], + cwd=root, + check=True, + capture_output=True, + text=True, + timeout=30, + ) + assert result.stdout == "final answer\n" diff --git a/tests/test_memory_recall.py b/tests/test_memory_recall.py index e86b6ab..30faa80 100644 --- a/tests/test_memory_recall.py +++ b/tests/test_memory_recall.py @@ -1,6 +1,7 @@ """Source-linked recall through public stores and tool dispatch.""" import json +import tracemalloc import pytest @@ -61,6 +62,45 @@ def test_source_snapshot_covers_messages_across_event_gaps(tmp_path): ) +def test_source_snapshot_memory_depends_on_selected_range(tmp_path): + sessions = SessionStore(tmp_path / "cfg") + session = sessions.create("source", tmp_path) + sessions.append_message(session, {"role": "user", "content": "selected evidence"}) + for _ in range(128): + sessions.append_message(session, {"role": "assistant", "content": "x" * 65536}) + + tracemalloc.start() + try: + snapshot = sessions.source_snapshot(session.id, 1, 1, project_root=tmp_path) + _, peak = tracemalloc.get_traced_memory() + finally: + tracemalloc.stop() + assert snapshot.status == "valid" + assert snapshot.messages == ({"role": "user", "content": "selected evidence"},) + # Reading an 8 MiB archive must not retain its unrelated message payloads. + assert peak < 2 * 1024 * 1024 + + +@pytest.mark.parametrize("change", ["corrupt", "duplicate", "out_of_order", "coerced_seq"]) +def test_source_snapshot_validates_records_outside_selected_range(tmp_path, change): + sessions = SessionStore(tmp_path / "cfg") + session = sessions.create("source", tmp_path) + for _ in range(3): + sessions.append_message(session, {"role": "user", "content": "evidence"}) + records = [json.loads(line) for line in session.path.read_text().splitlines()] + if change == "duplicate": + records[-1]["seq"] = 2 + elif change == "out_of_order": + records[-2:] = records[-2:][::-1] + elif change == "coerced_seq": + records[-1]["seq"] = "3" + session.path.write_text("\n".join(json.dumps(record) for record in records) + "\n") + if change == "corrupt": + with session.path.open("a") as stream: + stream.write("broken record\n") + assert sessions.source_snapshot(session.id, 1, 1, project_root=tmp_path).status == "stale" + + def test_fact_source_attachment_persists_and_search_is_bounded(tmp_path): sessions = SessionStore(tmp_path / "cfg") session = sessions.create("source", tmp_path) diff --git a/tests/test_proc.py b/tests/test_proc.py index d3a5ec0..cd0bb05 100644 --- a/tests/test_proc.py +++ b/tests/test_proc.py @@ -5,8 +5,11 @@ import asyncio import contextlib import sys +import tracemalloc -from lecode.extras.proc import run_proc +import pytest + +from lecode.extras.proc import TRUNCATION_MARKER, run_proc async def test_echo_round_trip(): @@ -47,6 +50,70 @@ async def test_output_cap_head_tail(): assert result.stdout.rstrip().endswith("x") +@pytest.mark.parametrize("cap", [1, 2, 3, 32, 1000, 10000]) +async def test_streamed_cap_preserves_byte_boundaries(cap): + from lecode.extras.proc import _read_capped + + data = b"\xff" + "é🙂".encode() * 1000 + b"\xf0" + stream = asyncio.StreamReader() + stream.feed_data(data) + stream.feed_eof() + text, truncated = await _read_capped(stream, cap) + assert truncated == (len(data) > cap) + if truncated: + head = data[: cap // 2] + tail = data[-(cap // 2) :] if cap // 2 else b"" + assert text == ( + head.decode("utf-8", errors="replace") + + TRUNCATION_MARKER.format(skipped=len(data) - len(head) - len(tail)) + + tail.decode("utf-8", errors="replace") + ) + else: + assert text == data.decode("utf-8", errors="replace") + + +async def test_large_stdout_and_stderr_are_drained_while_feeding_stdin(): + result = await run_proc( + [ + sys.executable, + "-c", + ( + "import sys\n" + "for _ in range(64):\n" + " sys.stdout.buffer.write(b'o'*65536)\n" + " sys.stderr.buffer.write(b'e'*65536)\n" + "print(len(sys.stdin.buffer.read()))\n" + ), + ], + input="i" * (2 * 1024 * 1024), + max_output_bytes=2048, + timeout=10, + ) + assert result.exit_code == 0 and result.truncated and not result.timed_out + assert result.stdout.startswith("o" * 1024) + assert result.stdout.endswith("2097152\n") + assert result.stderr.startswith("e" * 1024) and result.stderr.endswith("e" * 1024) + + +@pytest.mark.parametrize("cap", [1, 2048]) +async def test_output_retention_is_bounded(cap): + tracemalloc.start() + try: + result = await run_proc( + [ + sys.executable, + "-c", + ("import sys\nfor _ in range(128): sys.stdout.buffer.write(b'x'*65536)"), + ], + max_output_bytes=cap, + ) + _, peak = tracemalloc.get_traced_memory() + finally: + tracemalloc.stop() + assert result.exit_code == 0 and result.truncated + assert peak < 2 * 1024 * 1024 + + async def test_cwd(): result = await run_proc(["/bin/pwd"], cwd="/tmp") assert result.stdout.strip().endswith("tmp") diff --git a/tests/test_session_storage.py b/tests/test_session_storage.py index c8687b7..bf14bb0 100644 --- a/tests/test_session_storage.py +++ b/tests/test_session_storage.py @@ -4,6 +4,7 @@ import json import os +import tracemalloc import pytest @@ -205,11 +206,48 @@ def test_open_round_trip(store, session): assert reopened.next_seq == session.next_seq == 5 +@pytest.mark.parametrize("name", ["renamed", "", None]) +def test_open_preserves_latest_rename_max_sequence_and_corrupt_count(store, session, name): + store.append_event(session, "rename", {"name": "earlier rename"}) + store.append_event(session, "rename", {"name": name}) + with session.path.open("a") as stream: + stream.write("broken record\n") + for seq in (42, 2): + stream.write( + MessageRecord( + seq=seq, ts="timestamp", role="user", message={"role": "user", "content": "x"} + ).model_dump_json() + + "\n" + ) + reopened = store.open(session.id) + assert reopened.meta.name == (name or "demo") + assert reopened.next_seq == 43 + assert store.corrupt_lines == 1 + + def test_open_missing_raises(store): with pytest.raises(SessionNotFoundError): store.open("nope") +def test_load_grants_streams_unrelated_messages(store, session): + store.grant_permission(session, "read", "*.py") + for _ in range(128): + store.append_message(session, {"role": "assistant", "content": "x" * 65536}) + store.grant_permission(session, "bash", "git *") + with session.path.open("a") as stream: + stream.write("broken record\n") + tracemalloc.start() + try: + grants = store.load_grants(session) + _, peak = tracemalloc.get_traced_memory() + finally: + tracemalloc.stop() + assert grants == [("read", "*.py"), ("bash", "git *")] + assert store.corrupt_lines == 1 + assert peak < 2 * 1024 * 1024 + + def test_append_assigns_increasing_seq(store, session): record = store.append_message(session, {"role": "user", "content": "third"}) assert record.seq == 5 diff --git a/tests/test_setup_wizard.py b/tests/test_setup_wizard.py index a308f9c..e4defb5 100644 --- a/tests/test_setup_wizard.py +++ b/tests/test_setup_wizard.py @@ -293,7 +293,7 @@ async def fake_wizard(): called.append(True) return config_dir() / "config.toml" - monkeypatch.setattr("lecode.cli.run_wizard", fake_wizard) + monkeypatch.setattr("lecode.setup_wizard.run_wizard", fake_wizard) assert run_setup() == 0 assert called == [True] @@ -306,7 +306,7 @@ def test_setup_cancelled_exits_1(cfg_dir, monkeypatch): async def cancelled(): raise KeyboardInterrupt - monkeypatch.setattr("lecode.cli.run_wizard", cancelled) + monkeypatch.setattr("lecode.setup_wizard.run_wizard", cancelled) assert run_setup() == 1 @@ -352,8 +352,8 @@ async def fake_offer(): async def no_name(store, **kwargs): return None # Ctrl-D at the name prompt → exit 0 - monkeypatch.setattr("lecode.cli.offer_first_run_setup", fake_offer) - monkeypatch.setattr("lecode.cli.prompt_session_name", no_name) + monkeypatch.setattr("lecode.setup_wizard.offer_first_run_setup", fake_offer) + monkeypatch.setattr("lecode.tui.name_prompt.prompt_session_name", no_name) def test_interactive_first_run_offers_setup(tmp_path, monkeypatch): diff --git a/tests/test_tool_bash.py b/tests/test_tool_bash.py index 8cf20a0..b310a8e 100644 --- a/tests/test_tool_bash.py +++ b/tests/test_tool_bash.py @@ -5,9 +5,12 @@ import asyncio import contextlib import os +import shlex import stat import subprocess +import pytest + from lecode.agent.tools import bash from lecode.agent.tools.bash import MAX_OUTPUT_BYTES from lecode.extras import rtk @@ -64,6 +67,88 @@ async def test_truncation_and_overflow_file(tool_ctx, tmp_path): assert result.metadata["proc_result"].truncated +@pytest.mark.parametrize( + "payload", + [ + b"x" * (MAX_OUTPUT_BYTES - 1), + b"x" * MAX_OUTPUT_BYTES, + "🙂".encode() * 20000, + b"a" * 65535 + "🙂".encode() + b"\xff\r\n" * 30000 + b"\xf0", + ], + ids=["below_cap", "at_cap", "unicode", "invalid_utf8"], +) +async def test_output_and_overflow_preserve_decoding(tool_ctx, tmp_path, monkeypatch, payload): + async def unchanged(command): + return command + + monkeypatch.setattr(bash, "rewrite_command", unchanged) + source = tmp_path / "output.bin" + source.write_bytes(payload) + result = await bash.make_tool().run({"command": f"cat {shlex.quote(str(source))}"}, tool_ctx) + assert not result.is_error + text = payload.decode("utf-8", errors="replace") + files = list((tmp_path / "cfg" / "overflow").glob("*.log")) + if len(payload) > MAX_OUTPUT_BYTES: + assert len(files) == 1 + assert files[0].read_bytes() == text.encode() + half = MAX_OUTPUT_BYTES // 2 + expected = ( + text[:half] + + f"\n… [output truncated; full output saved to {files[0]}] …\n" + + text[-half:] + ) + else: + assert not files + expected = text + assert result.metadata["proc_result"].stdout == expected + assert result.content == expected.rstrip("\n") + assert result.metadata["proc_result"].truncated == (len(payload) > MAX_OUTPUT_BYTES) + + +@pytest.mark.parametrize("late_child", [False, True]) +async def test_output_write_failure_kills_child(tmp_path, monkeypatch, late_child): + kill_tree = bash._kill_tree + first_kill = True + + def miss_child_during_first_signal(proc): + nonlocal first_kill + if first_kill: + first_kill = False + proc.kill() # Reproduce a child missed while the shell is spawning it. + else: + kill_tree(proc) + + if late_child: + monkeypatch.setattr(bash, "_kill_tree", miss_child_during_first_signal) + + def failed_write(chunk): + raise OSError("output disk full") + + slot = {} + with pytest.raises(OSError, match="output disk full"): + await bash._run_shell( + "sleep 30 & echo ready; wait" if late_child else "echo ready; sleep 30", + tmp_path, + 30, + 30, + MAX_OUTPUT_BYTES, + on_chunk=failed_write, + proc_slot=slot, + ) + try: + assert slot["proc"].returncode is not None + remaining = subprocess.run( + ["pgrep", "-g", str(slot["proc"].pid)], capture_output=True, text=True + ).stdout.split() + assert not remaining, subprocess.run( + ["ps", "-p", ",".join(remaining), "-o", "pid,ppid,pgid,state,command"], + capture_output=True, + text=True, + ).stdout + finally: + kill_tree(slot["proc"]) + + async def test_rtk_rewrite_applied(tool_ctx, tmp_path, monkeypatch): # Fake rtk: `rtk rewrite "echo hello"` → "echo HELLO" (the rewritten # command is what actually runs). @@ -98,4 +183,8 @@ async def test_cancel_kills_child(tool_ctx): with contextlib.suppress(asyncio.CancelledError): await task out = subprocess.run(["pgrep", "-f", "sleep 30"], capture_output=True).stdout - assert out == b"" + assert out == b"", subprocess.run( + ["ps", "-p", ",".join(out.decode().split()), "-o", "pid,ppid,pgid,state,command"], + capture_output=True, + text=True, + ).stdout diff --git a/tests/test_tool_read.py b/tests/test_tool_read.py index 90535aa..5ab90af 100644 --- a/tests/test_tool_read.py +++ b/tests/test_tool_read.py @@ -2,6 +2,10 @@ from __future__ import annotations +from pathlib import Path + +import pytest + from lecode.agent.tools.read import line_anchor, make_tool @@ -74,3 +78,38 @@ def test_anchor_format(): anchor = line_anchor(12, " hello world ") assert anchor.startswith("12:") assert len(anchor.split(":")[1]) == 2 + + +@pytest.mark.parametrize( + "text,offset,limit", + [ + ("one\r\ntwo\rthree\n", 2, 1), + ("one\vtwo\fthree\x85four\u2028five\u2029\nlast", 2, 3), + ("\n\none\n\n", 1, 3), + ("one\ntwo", 10, 2), + ("", 1, 2), + ("one\ntwo\nthree\nfour", 1, -2), + ], +) +async def test_read_streams_page_with_splitlines_semantics( + tool_ctx, tmp_path, monkeypatch, text, offset, limit +): + path = tmp_path / "page.txt" + path.write_text(text, encoding="utf-8") + + def deny_whole_file_read(*args, **kwargs): + pytest.fail("paginated reads must not load the entire file") + + if limit > 0: + monkeypatch.setattr(Path, "read_text", deny_whole_file_read) + result = await make_tool().run( + {"path": "page.txt", "offset": offset, "limit": limit, "with_anchors": True}, tool_ctx + ) + lines = text.splitlines() + page = lines[offset - 1 : offset - 1 + limit] + expected = [f"{line_anchor(i, line)}\t{line}" for i, line in enumerate(page, offset)] + remaining = len(lines) - (offset - 1 + len(page)) + if remaining > 0: + expected.append(f"… {remaining} more lines (continue with offset={offset + len(page)})") + assert result.content == "\n".join(expected or ["(empty file)"]) + assert str(path) in tool_ctx.read_paths diff --git a/tests/test_tui_app.py b/tests/test_tui_app.py index 5c55636..d6f5af7 100644 --- a/tests/test_tui_app.py +++ b/tests/test_tui_app.py @@ -1322,7 +1322,7 @@ def cli_env(tmp_path, monkeypatch): monkeypatch.setattr("lecode.cli.check_dependencies", lambda: None) monkeypatch.setattr("lecode.cli.build_provider", lambda config, api_key=None: object()) FakeTui.instances = [] - monkeypatch.setattr("lecode.cli.TuiApp", FakeTui) + monkeypatch.setattr("lecode.tui.app.TuiApp", FakeTui) return tmp_path @@ -1334,7 +1334,7 @@ async def _prompt(store, **kwargs): def test_cli_abort_exits_zero_without_session(cli_env, monkeypatch): - monkeypatch.setattr("lecode.cli.prompt_session_name", _name_prompt(None)) + monkeypatch.setattr("lecode.tui.name_prompt.prompt_session_name", _name_prompt(None)) result = runner.invoke(cli_app, []) assert result.exit_code == 0 assert FakeTui.instances == [] @@ -1342,7 +1342,7 @@ def test_cli_abort_exits_zero_without_session(cli_env, monkeypatch): def test_cli_interactive_creates_named_session(cli_env, monkeypatch): - monkeypatch.setattr("lecode.cli.prompt_session_name", _name_prompt("chatty")) + monkeypatch.setattr("lecode.tui.name_prompt.prompt_session_name", _name_prompt("chatty")) result = runner.invoke(cli_app, []) assert result.exit_code == 0 assert len(FakeTui.instances) == 1 @@ -1351,21 +1351,21 @@ def test_cli_interactive_creates_named_session(cli_env, monkeypatch): def test_cli_default_mode_is_yolo(cli_env, monkeypatch): - monkeypatch.setattr("lecode.cli.prompt_session_name", _name_prompt("s")) + monkeypatch.setattr("lecode.tui.name_prompt.prompt_session_name", _name_prompt("s")) result = runner.invoke(cli_app, []) assert result.exit_code == 0 assert FakeTui.instances[0].runtime.ctx.permission_checker.mode == "yolo" def test_cli_safe_flag_forces_readonly(cli_env, monkeypatch): - monkeypatch.setattr("lecode.cli.prompt_session_name", _name_prompt("s")) + monkeypatch.setattr("lecode.tui.name_prompt.prompt_session_name", _name_prompt("s")) result = runner.invoke(cli_app, ["--safe"]) assert result.exit_code == 0 assert FakeTui.instances[0].runtime.ctx.permission_checker.mode == "readonly" def test_cli_read_only_alias_still_works(cli_env, monkeypatch): - monkeypatch.setattr("lecode.cli.prompt_session_name", _name_prompt("s")) + monkeypatch.setattr("lecode.tui.name_prompt.prompt_session_name", _name_prompt("s")) result = runner.invoke(cli_app, ["--read-only"]) assert result.exit_code == 0 assert FakeTui.instances[0].runtime.ctx.permission_checker.mode == "readonly" @@ -1377,7 +1377,7 @@ def test_cli_resume_keeps_name_without_prompt(cli_env, monkeypatch): async def _boom(store, **kwargs): raise AssertionError("name prompt must not run on --resume") - monkeypatch.setattr("lecode.cli.prompt_session_name", _boom) + monkeypatch.setattr("lecode.tui.name_prompt.prompt_session_name", _boom) result = runner.invoke(cli_app, ["-r", "old-session"]) assert result.exit_code == 0 assert FakeTui.instances[0].session.name == "old-session" @@ -1387,14 +1387,14 @@ def test_cli_continue_picks_latest(cli_env, monkeypatch): store = SessionStore() store.create("first", cli_env) store.create("second", cli_env) - monkeypatch.setattr("lecode.cli.prompt_session_name", _name_prompt(None)) + monkeypatch.setattr("lecode.tui.name_prompt.prompt_session_name", _name_prompt(None)) result = runner.invoke(cli_app, ["-c"]) assert result.exit_code == 0 assert FakeTui.instances[0].session.name == "second" def test_cli_resume_unknown_ref_fails(cli_env, monkeypatch): - monkeypatch.setattr("lecode.cli.prompt_session_name", _name_prompt(None)) + monkeypatch.setattr("lecode.tui.name_prompt.prompt_session_name", _name_prompt(None)) result = runner.invoke(cli_app, ["-r", "nope"]) assert result.exit_code == 2 assert "nope" in result.output @@ -1426,7 +1426,7 @@ async def _pick(store, cwd, **kwargs): def test_cli_no_color_lands_in_config(cli_env, monkeypatch): - monkeypatch.setattr("lecode.cli.prompt_session_name", _name_prompt("x")) + monkeypatch.setattr("lecode.tui.name_prompt.prompt_session_name", _name_prompt("x")) result = runner.invoke(cli_app, ["--no-color"]) assert result.exit_code == 0 assert FakeTui.instances[0].config.ui.no_color is True @@ -1547,7 +1547,7 @@ def test_cli_resume_locked_session_fails(cli_env, monkeypatch): session = store.create("busy", cli_env) lock = store.acquire_lock(session) assert lock is not None - monkeypatch.setattr("lecode.cli.prompt_session_name", _name_prompt(None)) + monkeypatch.setattr("lecode.tui.name_prompt.prompt_session_name", _name_prompt(None)) result = runner.invoke(cli_app, ["-r", "busy"]) assert result.exit_code == 2 assert "already open in another lecode process" in result.output diff --git a/tests/test_tui_loading.py b/tests/test_tui_loading.py index d40d4e6..54a78b8 100644 --- a/tests/test_tui_loading.py +++ b/tests/test_tui_loading.py @@ -344,9 +344,9 @@ def test_interactive_startup_prints_loading_screen(env, monkeypatch, capsys): """run_interactive prints the banner and step lines before the chat.""" import lecode.cli as cli - monkeypatch.setattr(cli, "prompt_session_name", _fake_name_prompt) + monkeypatch.setattr("lecode.tui.name_prompt.prompt_session_name", _fake_name_prompt) monkeypatch.setattr(cli, "build_provider", lambda config, api_key=None: object()) - monkeypatch.setattr(cli, "TuiApp", _FakeTui) + monkeypatch.setattr("lecode.tui.app.TuiApp", _FakeTui) monkeypatch.setattr(cli, "_run_tui", _fake_run_tui) code = cli.run_interactive() assert code == 0 @@ -379,7 +379,7 @@ def test_catalog_fetch_runs_in_background(env, monkeypatch): from lecode.providers.live import LoadedCatalog events: list[str] = [] - monkeypatch.setattr(cli, "prompt_session_name", _fake_name_prompt) + monkeypatch.setattr("lecode.tui.name_prompt.prompt_session_name", _fake_name_prompt) monkeypatch.setattr(cli, "build_provider", lambda config, api_key=None: object()) class FakeTui: @@ -400,7 +400,7 @@ def slow_fetch(config, api_key=None): events.append("fetch-done") return LoadedCatalog(Catalog.default(), "live", 427) - monkeypatch.setattr(cli, "TuiApp", FakeTui) + monkeypatch.setattr("lecode.tui.app.TuiApp", FakeTui) monkeypatch.setattr(cli, "_run_tui", fake_run_tui) monkeypatch.setattr(cli, "fetch_catalog", slow_fetch) assert cli.run_interactive() == 0 diff --git a/tests/test_worktree.py b/tests/test_worktree.py index a2f7049..0d76f96 100644 --- a/tests/test_worktree.py +++ b/tests/test_worktree.py @@ -1051,13 +1051,13 @@ def cli_env(tmp_path, monkeypatch): monkeypatch.setattr("lecode.cli.check_dependencies", lambda: None) monkeypatch.setattr("lecode.cli.build_provider", lambda config, api_key=None: object()) FakeTui.instances = [] - monkeypatch.setattr("lecode.cli.TuiApp", FakeTui) + monkeypatch.setattr("lecode.tui.app.TuiApp", FakeTui) return tmp_path def test_cli_worktree_flag_switches_cwd(cli_env, monkeypatch): make_repo_sync(cli_env) - monkeypatch.setattr("lecode.cli.prompt_session_name", _name_prompt("wt-session")) + monkeypatch.setattr("lecode.tui.name_prompt.prompt_session_name", _name_prompt("wt-session")) result = runner.invoke(cli_app, ["--worktree", "feat"]) assert result.exit_code == 0, result.output expected = cli_env / ".lecode" / "worktrees" / "feat" @@ -1072,7 +1072,7 @@ def test_cli_worktree_flag_switches_cwd(cli_env, monkeypatch): def test_cli_worktree_flag_not_a_repo(cli_env, monkeypatch): - monkeypatch.setattr("lecode.cli.prompt_session_name", _name_prompt("x")) + monkeypatch.setattr("lecode.tui.name_prompt.prompt_session_name", _name_prompt("x")) result = runner.invoke(cli_app, ["--worktree", "feat"]) assert result.exit_code == 2 assert "not a git repository" in result.output