diff --git a/src/loushang/harness/session/inspection.py b/src/loushang/harness/session/inspection.py index 9b392d083..3dbe048a1 100644 --- a/src/loushang/harness/session/inspection.py +++ b/src/loushang/harness/session/inspection.py @@ -23,6 +23,7 @@ ) from loushang.harness.transcript import ( AGENT_MESSAGE_KIND, + CONTEXT_COMPACTION_CHECKPOINT_KIND, MODEL_CALL_OUTCOME_KIND, MODEL_INPUT_PREPARED_KIND, AgentTranscriptInspector, @@ -222,7 +223,7 @@ def get_context_usage(self) -> ContextUsage: branch_entries: list[object] = list(branch_records) transcript_revision = len(self.session.get_entries()) leaf_id = self.session.get_leaf_id() - counts = self._transcript.message_counts() + counts = self._transcript.message_counts(messages=messages) snapshot = build_context_usage_snapshot( messages, branch_entries, @@ -371,6 +372,14 @@ def _derive_provider_context_anchor( self, branch_entries: Sequence[AgentTranscriptRecord], ) -> ProviderContextAnchor | None: + # A provider observation made before compaction describes the old full + # prompt. Rebuilding it is expensive and cannot calibrate the summary. + for index in range(len(branch_entries) - 1, -1, -1): + if branch_entries[index].kind == CONTEXT_COMPACTION_CHECKPOINT_KIND: + branch_entries = branch_entries[index + 1 :] + if not branch_entries: + return None + break ledger = project_model_call_usage(tuple(branch_entries)) for attempt in reversed(ledger.attempts): if not attempt.terminal: diff --git a/src/loushang/harness/transcript/interaction.py b/src/loushang/harness/transcript/interaction.py index 77640a4cc..ed1df8a3f 100644 --- a/src/loushang/harness/transcript/interaction.py +++ b/src/loushang/harness/transcript/interaction.py @@ -13,7 +13,7 @@ from dataclasses import dataclass, field from typing import Protocol, cast -from loushang.agent.types import ThinkingLevel +from loushang.agent.types import AgentMessage, ThinkingLevel from loushang.ai.model import Model, ModelSelection from loushang.ai.types import ( AssistantMessage, @@ -136,13 +136,16 @@ class AgentTranscriptInspector: session: AgentTranscriptSession - def message_counts(self) -> TranscriptMessageCounts: - context = self.session.build_context() + def message_counts( + self, *, messages: Sequence[AgentMessage] | None = None + ) -> TranscriptMessageCounts: + if messages is None: + messages = self.session.build_context().messages assistant_message_count = 0 user_message_count = 0 tool_call_count = 0 tool_result_count = 0 - for message in context.messages: + for message in messages: if isinstance(message, AssistantMessage): assistant_message_count += 1 tool_call_count += sum( @@ -153,7 +156,7 @@ def message_counts(self) -> TranscriptMessageCounts: elif isinstance(message, ToolResultMessage): tool_result_count += 1 return TranscriptMessageCounts( - message_count=len(context.messages), + message_count=len(messages), assistant_message_count=assistant_message_count, user_message_count=user_message_count, tool_call_count=tool_call_count, diff --git a/src/loushang/harnesstui/conversation/screen_app.py b/src/loushang/harnesstui/conversation/screen_app.py index e06e70ba6..d9b41a9e1 100644 --- a/src/loushang/harnesstui/conversation/screen_app.py +++ b/src/loushang/harnesstui/conversation/screen_app.py @@ -73,15 +73,15 @@ class ScreenConversationApp: now: Callable[[], float] = time.monotonic composer: Composer = field(default_factory=Composer) state: ScreenConversationState = field(init=False) - capability_provider: Callable[[], ConversationCapabilities] | None = field(default=None, kw_only=True) + capability_provider: Callable[[], ConversationCapabilities] | None = field( + default=None, kw_only=True + ) active_surface: Any | None = None surface_host: SurfaceHost | None = None transcript_theme: ThemeResolver | None = None welcome_theme: ThemeResolver | None = None active_transcript_line_budget: int = 0 - compaction_summary_formatter: Callable[[str], str] = ( - _normalized_compaction_summary - ) + compaction_summary_formatter: Callable[[str], str] = _normalized_compaction_summary stable_render_cache_entry_limit: int = DEFAULT_STABLE_TRANSCRIPT_CACHE_ENTRY_LIMIT render_requester: Callable[[RenderRequestKind], object] | None = None terminal_diagnostics_provider: Callable[[], str] | None = None @@ -184,6 +184,7 @@ def end_assistant(self, final_text: str | None = None) -> None: def complete_run(self, *, elapsed_seconds: float | None = None) -> None: elapsed = self.elapsed_seconds() if elapsed_seconds is None else elapsed_seconds self.state.complete_run(elapsed_seconds=elapsed) + self._context_usage_refresh_key = None self._transcript_region.clear_transient_cache() def queue_followup(self, text: str) -> None: @@ -406,7 +407,9 @@ def render(self, constraints: RenderConstraints) -> RenderResult: return layout.render(constraints) def _refresh_context_usage(self) -> None: - if self.context_usage_provider is None: + # Rebuilding a long session's context is synchronous and can stall + # animation and input while a tool is running. Refresh after the run. + if self.context_usage_provider is None or self.state.running: return refresh_key = ( self.state.records_revision, diff --git a/src/loushang/tui/terminal_input.py b/src/loushang/tui/terminal_input.py index 2ad88a7c9..57a369411 100644 --- a/src/loushang/tui/terminal_input.py +++ b/src/loushang/tui/terminal_input.py @@ -159,14 +159,6 @@ async def read_input_chunk_or_render_tick( if active_task is not None and active_task.done(): return None - wait_for: set[asyncio.Task[Any]] = {input_task} - if active_task is not None and not active_task.done(): - wait_for.add(active_task) - render_task: asyncio.Task[bool] | None = None - if render_wakeup is not None: - render_task = asyncio.create_task(render_wakeup.wait()) - wait_for.add(render_task) - decision = runtime.request_next_animation_frame() timeout = None timeout_reason = "render" @@ -186,16 +178,27 @@ async def read_input_chunk_or_render_tick( timeout = idle_timeout timeout_reason = "idle_wakeup" - done, _pending = await asyncio.wait( - wait_for, timeout=timeout, return_when=asyncio.FIRST_COMPLETED - ) + wait_for: set[asyncio.Task[Any]] = {input_task} + if active_task is not None and not active_task.done(): + wait_for.add(active_task) + render_task: asyncio.Task[bool] | None = None + if render_wakeup is not None: + render_task = asyncio.create_task(render_wakeup.wait()) + wait_for.add(render_task) + + try: + done, _pending = await asyncio.wait( + wait_for, timeout=timeout, return_when=asyncio.FIRST_COMPLETED + ) + finally: + if render_task is not None: + if not render_task.done(): + render_task.cancel() + with suppress(asyncio.CancelledError): + await render_task render_wakeup_fired = render_task is not None and render_task in done if render_wakeup_fired and render_wakeup is not None: render_wakeup.clear() - if render_task is not None and not render_task.done(): - render_task.cancel() - with suppress(asyncio.CancelledError): - await render_task if input_task in done: return input_task.result() if active_task is not None and active_task in done: @@ -245,6 +248,8 @@ def _read_input_chunk_blocking(stdin: Any) -> str: def stream_is_tty(stream: Any) -> bool: isatty = getattr(stream, "isatty", None) return bool(callable(isatty) and isatty()) + + __all__ = [ "ESCAPE_SEQUENCE_IDLE_TIMEOUT_MS", "TerminalInputMode", diff --git a/tests/harness/session/test_inspection.py b/tests/harness/session/test_inspection.py index 50e16bea2..43d2f3dd7 100644 --- a/tests/harness/session/test_inspection.py +++ b/tests/harness/session/test_inspection.py @@ -26,6 +26,7 @@ from loushang.harness.session.inspection import _build_token_usage_totals from loushang.harness.transcript import ( AGENT_MESSAGE_KIND, + CONTEXT_COMPACTION_CHECKPOINT_KIND, MODEL_CALL_ATTEMPT_USAGE_KIND, MODEL_CALL_OUTCOME_KIND, MODEL_INPUT_PREPARED_KIND, @@ -200,6 +201,44 @@ def test_provider_anchor_surface_requires_the_current_model_projection() -> None assert inspector._measure_current_model_surface(messages) is None +def test_provider_anchor_does_not_rebuild_input_before_compaction(monkeypatch) -> None: + inspector = asyncio.run(_inspector()) + records = [ + _record("old-snapshot", MODEL_INPUT_PREPARED_KIND, _snapshot("old", 1)), + _record( + "old-usage", + MODEL_CALL_ATTEMPT_USAGE_KIND, + _attempt_usage("old", 1, input=10, terminal=True), + ), + _record("checkpoint", CONTEXT_COMPACTION_CHECKPOINT_KIND, object()), + ] + + def reject_rebuild(snapshot_id: str) -> None: + raise AssertionError(f"rebuilt compacted input {snapshot_id}") + + monkeypatch.setattr(inspector.session, "rebuild_model_input", reject_rebuild) + + assert inspector._derive_provider_context_anchor(records) is None + + +def test_context_usage_reuses_its_replayed_messages_for_counts(monkeypatch) -> None: + inspector = asyncio.run(_inspector()) + build_context = inspector.session.build_context + replay_count = 0 + + def count_replays(): + nonlocal replay_count + replay_count += 1 + return build_context() + + monkeypatch.setattr(inspector.session, "build_context", count_replays) + + usage = inspector.get_context_usage() + + assert usage.message_count == 3 + assert replay_count == 1 + + def test_token_totals_prefer_outcomes_and_keep_uncovered_legacy_usage() -> None: legacy_usage = Usage(1, 2, 3, 4, 10, None) outcome_usage = Usage(10, 5, 2, 1, 18, None) diff --git a/tests/harnesstui/conversation/test_screen_app.py b/tests/harnesstui/conversation/test_screen_app.py index 46fd72d7b..4b4199825 100644 --- a/tests/harnesstui/conversation/test_screen_app.py +++ b/tests/harnesstui/conversation/test_screen_app.py @@ -11,7 +11,9 @@ from loushang.tui.transcript import ( AssistantMessageRecord, ContextCompactionRecord, + ToolExecutionRecord, UserPromptRecord, + WorkedDividerRecord, ) from loushang.tui.ui_parts.transcript import TranscriptRegion @@ -71,6 +73,67 @@ def test_screen_conversation_app_keeps_presentation_and_region_instances() -> No assert app._transcript_region is region +@pytest.mark.tui_render_contract +@pytest.mark.parametrize("start_method", ("prompt", "begin_run")) +def test_active_frames_defer_context_usage_rebuild_until_run_completes( + start_method: str, +) -> None: + app = _app() + calls: list[int] = [] + + def context_usage() -> int: + calls.append(app.state.records_revision) + return len(calls) + + app.context_usage_provider = context_usage + constraints = RenderConstraints(width=60, max_height=24, visible_height=24) + app.render(constraints) + assert calls == [0] + + if start_method == "prompt": + app.start_prompt("run", started_at=1.0) + else: + app.begin_run(started_at=1.0) + app.render(constraints) + app.state.upsert_tool_record( + "tool-1", ToolExecutionRecord(name="bash", state="running", elapsed_seconds=0.0) + ) + app.render(constraints) + + assert calls == [0] + assert app.state.context_usage == 1 + + app.complete_run(elapsed_seconds=1.0) + app.render(constraints) + + assert calls == [0, app.state.records_revision] + assert app.state.context_usage == 2 + + +def test_context_usage_refreshes_after_run_without_display_changes() -> None: + app = _app() + app.state.records.append(WorkedDividerRecord(1.0)) + app.state.mark_records_changed() + calls = 0 + + def context_usage() -> int: + nonlocal calls + calls += 1 + return calls + + app.context_usage_provider = context_usage + app.statusline_preview_snapshot() + revision = app.state.records_revision + + app.begin_run(started_at=1.0) + app.statusline_preview_snapshot() + app.complete_run(elapsed_seconds=1.0) + app.statusline_preview_snapshot() + + assert app.state.records_revision == revision + assert calls == 2 + + def test_screen_conversation_app_reports_window_replacement_reason_once() -> None: app = _app() @@ -82,9 +145,7 @@ def test_screen_conversation_app_reports_window_replacement_reason_once() -> Non assert app.consume_render_baseline_reset_reason() is None -def test_screen_conversation_app_atomically_installs_bounded_resumed_history() -> ( - None -): +def test_screen_conversation_app_atomically_installs_bounded_resumed_history() -> None: app = _app() app.active_transcript_line_budget = 2 diff --git a/tests/tui/test_terminal_input.py b/tests/tui/test_terminal_input.py index ea22b7d2e..3a0d5abea 100644 --- a/tests/tui/test_terminal_input.py +++ b/tests/tui/test_terminal_input.py @@ -44,9 +44,7 @@ def test_terminal_input_mode_holds_native_lease_until_protocol_cleanup( "loushang.tui.terminal_input.drain_input", lambda *args, **kwargs: "" ) - with TerminalInputMode( - stdin=_TtyInput(), stdout=stdout, keyboard_protocols=False - ): + with TerminalInputMode(stdin=_TtyInput(), stdout=stdout, keyboard_protocols=False): assert lease.restore_calls == 0 assert factory.open_calls == 1 @@ -128,6 +126,86 @@ async def wake_later() -> None: assert decisions >= 2 +def test_render_waiter_is_cleaned_after_immediate_render() -> None: + async def run() -> tuple[str | None, int]: + release_input = asyncio.Event() + + class ImmediateOnceRuntime: + decisions = 0 + + def request_next_animation_frame(self): + self.decisions += 1 + return ( + _ImmediateDecision() if self.decisions == 1 else _DelayedDecision() + ) + + def render_now(self) -> None: + release_input.set() + + async def read_after_render(_stdin: object) -> str: + await release_input.wait() + return "x" + + existing = asyncio.all_tasks() + result = await read_input_chunk_or_render_tick( + StringIO(), + runtime=ImmediateOnceRuntime(), + active_task=None, + input_chunk_reader=read_after_render, + render_wakeup=asyncio.Event(), + ) + pending = asyncio.all_tasks() - existing + try: + return result, len(pending) + finally: + for task in pending: + task.cancel() + await asyncio.gather(*pending, return_exceptions=True) + + assert asyncio.run(run()) == ("x", 0) + + +def test_render_waiter_is_cleaned_when_input_wait_is_cancelled() -> None: + async def run() -> int: + waiting = asyncio.Event() + + class WaitingRuntime: + def request_next_animation_frame(self): + waiting.set() + return _DelayedDecision() + + def render_now(self) -> None: + pass + + async def blocked_read(_stdin: object) -> str: + await asyncio.Event().wait() + return "x" + + existing = asyncio.all_tasks() + reader = asyncio.create_task( + read_input_chunk_or_render_tick( + StringIO(), + runtime=WaitingRuntime(), + active_task=None, + input_chunk_reader=blocked_read, + render_wakeup=asyncio.Event(), + ) + ) + await waiting.wait() + await asyncio.sleep(0) + reader.cancel() + await asyncio.gather(reader, return_exceptions=True) + pending = asyncio.all_tasks() - existing + try: + return len(pending) + finally: + for task in pending: + task.cancel() + await asyncio.gather(*pending, return_exceptions=True) + + assert asyncio.run(run()) == 0 + + def test_read_input_chunk_or_render_tick_wakes_for_terminal_runtime_deadline() -> None: async def run() -> tuple[str | None, int]: runtime = _Runtime()