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
49 changes: 47 additions & 2 deletions agent/context_compressor.py
Original file line number Diff line number Diff line change
Expand Up @@ -650,6 +650,10 @@ def __init__(
self.last_prompt_tokens = 0
self.last_completion_tokens = 0
self.last_real_prompt_tokens = 0
# How many conversation messages ``last_prompt_tokens`` accounts for.
# Lets ``live_context_tokens`` price only the tail appended since the
# last API call instead of freezing or re-estimating the whole list.
self.last_prompt_messages_len: Optional[int] = None
self.last_compression_rough_tokens = 0
self.last_rough_tokens_when_real_prompt_fit = 0
self.awaiting_real_usage_after_compression = False
Expand Down Expand Up @@ -681,9 +685,16 @@ def __init__(
self._last_aux_model_failure_error: Optional[str] = None
self._last_aux_model_failure_model: Optional[str] = None

def update_from_response(self, usage: Dict[str, Any]):
"""Update tracked token usage from API response."""
def update_from_response(self, usage: Dict[str, Any], messages_len: Optional[int] = None):
"""Update tracked token usage from API response.

``messages_len`` is the length of the conversation list the prompt was
built from, recorded so mid-turn estimates know where the
provider-counted prefix ends.
"""
self.last_prompt_tokens = usage.get("prompt_tokens", 0)
if messages_len is not None:
self.last_prompt_messages_len = messages_len
self.last_completion_tokens = usage.get("completion_tokens", 0)
self.last_total_tokens = usage.get("total_tokens", self.last_prompt_tokens + self.last_completion_tokens)
if self.last_prompt_tokens > 0:
Expand All @@ -695,6 +706,36 @@ def update_from_response(self, usage: Dict[str, Any]):
self.last_rough_tokens_when_real_prompt_fit = 0
self.awaiting_real_usage_after_compression = False

def live_context_tokens(self, messages: List[Dict[str, Any]]) -> int:
"""Best-effort size of the context *right now*, mid-turn.

``last_prompt_tokens`` is exact but frozen at the last API call (or
preflight estimate); assistant text and tool results appended since
aren't in it, so consumers that poll it mid-turn β€” the desktop context
bar during a long tool batch β€” appear stuck. Price the tail appended
since the snapshot on top of the known base instead.
"""
base = self.last_prompt_tokens or 0
snap = self.last_prompt_messages_len
if base > 0 and snap is not None and 0 <= snap <= len(messages):
tail = 0
if snap < len(messages):
try:
tail = estimate_messages_tokens_rough(messages[snap:])
except Exception:
tail = 0
return base + tail
try:
rough = estimate_messages_tokens_rough(messages)
except Exception:
rough = 0
if self.awaiting_real_usage_after_compression:
# Right after compression the stale pre-compression base would
# overstate a freshly shrunk conversation; trust the rough
# estimate of the compressed list until real usage arrives.
return rough or base
return max(base, rough)

def should_defer_preflight_to_real_usage(self, rough_tokens: int) -> bool:
"""Return True when a high rough preflight estimate is known-noisy.

Expand Down Expand Up @@ -2075,4 +2116,8 @@ def compress(self, messages: List[Dict[str, Any]], current_tokens: int = None, f
)
logger.info("Compression #%d complete", self.compression_count)

# The conversation list is being replaced wholesale; the old
# messages-length snapshot no longer maps onto the new list.
self.last_prompt_messages_len = None

return compressed
27 changes: 23 additions & 4 deletions agent/conversation_loop.py
Original file line number Diff line number Diff line change
Expand Up @@ -117,7 +117,9 @@ def _ra():
return run_agent


def _emit_preflight_token_usage(agent: Any, request_tokens: int) -> None:
def _emit_preflight_token_usage(
agent: Any, request_tokens: int, messages_len: int | None = None
) -> None:
"""Send the current request-size estimate through the live usage channel."""
if request_tokens <= 0:
return
Expand All @@ -129,6 +131,10 @@ def _emit_preflight_token_usage(agent: Any, request_tokens: int) -> None:
previous = getattr(compressor, "last_prompt_tokens", 0) or 0
if request_tokens > previous:
compressor.last_prompt_tokens = request_tokens
if messages_len is not None:
# Record which conversation prefix the seeded estimate covers
# so live_context_tokens prices only the tail appended later.
compressor.last_prompt_messages_len = messages_len
except Exception:
logger.debug("could not update preflight context estimate", exc_info=True)

Expand Down Expand Up @@ -678,6 +684,7 @@ def run_conversation(
# would re-introduce the very desync we're avoiding.
if _preflight_tokens > (_compressor.last_prompt_tokens or 0):
_compressor.last_prompt_tokens = _preflight_tokens
_compressor.last_prompt_messages_len = len(messages)

if _preflight_deferred:
logger.info(
Expand All @@ -700,6 +707,7 @@ def run_conversation(
f">= {_compressor.threshold_tokens:,} threshold. "
"This may take a moment."
)
_pre_compress_preflight_tokens = _preflight_tokens
# May need multiple passes for very large sessions with small
# context windows (each pass summarises the middle N turns).
for _pass in range(3):
Expand Down Expand Up @@ -739,6 +747,11 @@ def run_conversation(
_preflight_tokens = _post_preflight_tokens
if not _compressor.should_compress(_preflight_tokens):
break # Under threshold or anti-thrash guard stopped it
if _preflight_tokens < _pre_compress_preflight_tokens:
agent._emit_status(
f"πŸ“¦ Compression complete: ~{_pre_compress_preflight_tokens:,} "
f"β†’ ~{_preflight_tokens:,} tokens."
)

# Plugin hook: pre_llm_call
# Fired once per turn before the tool-calling loop. Plugins can
Expand Down Expand Up @@ -1143,7 +1156,9 @@ def run_conversation(
approx_request_tokens = estimate_request_tokens_rough(
api_messages, tools=agent.tools or None
)
_emit_preflight_token_usage(agent, approx_request_tokens)
_emit_preflight_token_usage(
agent, approx_request_tokens, messages_len=len(messages)
)

_runtime_context_error = _ollama_context_limit_error(
agent, approx_request_tokens
Expand Down Expand Up @@ -1995,7 +2010,9 @@ def _perform_api_call(next_api_kwargs):
"cache_write_tokens": canonical_usage.cache_write_tokens,
"reasoning_tokens": canonical_usage.reasoning_tokens,
}
agent.context_compressor.update_from_response(usage_dict)
agent.context_compressor.update_from_response(
usage_dict, messages_len=len(messages)
)

# Cache discovered context length after successful call.
# Only persist limits confirmed by the provider (parsed
Expand Down Expand Up @@ -3278,7 +3295,9 @@ def _perform_api_call(next_api_kwargs):
f"πŸ—œοΈ Compressed request estimate "
f"~{pre_compress_request_tokens:,} β†’ ~{post_compress_request_tokens:,} tokens, retrying..."
)
_emit_preflight_token_usage(agent, post_compress_request_tokens)
_emit_preflight_token_usage(
agent, post_compress_request_tokens, messages_len=len(messages)
)
time.sleep(2) # Brief pause between compression retries
restart_with_compressed_messages = True
break
Expand Down
93 changes: 33 additions & 60 deletions agent/tool_executor.py
Original file line number Diff line number Diff line change
Expand Up @@ -240,6 +240,35 @@ def _execute(next_args: dict) -> Any:
return result, observed_args


def _context_usage_for_tool_events(agent, messages: list) -> dict:
"""Build the usage payload attached to tool.start/complete events.

Context numbers come from ``ContextCompressor.live_context_tokens`` so the
UI context bar keeps moving while tool results pile up mid-turn instead of
freezing at the last API call's prompt size.
"""
usage = {
"model": getattr(agent, "model", ""),
"input": getattr(agent, "session_input_tokens", 0) or 0,
"output": getattr(agent, "session_output_tokens", 0) or 0,
"total": getattr(agent, "session_total_tokens", 0) or 0,
"calls": getattr(agent, "session_api_calls", 0) or 0,
}
comp = getattr(agent, "context_compressor", None)
if comp:
try:
ctx_used = comp.live_context_tokens(messages)
except Exception:
ctx_used = getattr(comp, "last_prompt_tokens", 0) or 0
ctx_used = ctx_used or usage["total"] or 0
ctx_max = getattr(comp, "context_length", 0) or 0
if ctx_max:
usage["context_used"] = ctx_used
usage["context_max"] = ctx_max
usage["context_percent"] = max(0, min(100, round(ctx_used / ctx_max * 100)))
return usage


def execute_tool_calls_concurrent(agent, assistant_message, messages: list, effective_task_id: str, api_call_count: int = 0) -> None:
"""Execute multiple tool calls concurrently using a thread pool.

Expand Down Expand Up @@ -440,21 +469,7 @@ def execute_tool_calls_concurrent(agent, assistant_message, messages: list, effe

# Inject context usage into tool events so frontend can update
# the context bar in real-time during tool execution.
usage = {
"model": getattr(agent, "model", ""),
"input": getattr(agent, "session_input_tokens", 0) or 0,
"output": getattr(agent, "session_output_tokens", 0) or 0,
"total": getattr(agent, "session_total_tokens", 0) or 0,
"calls": getattr(agent, "session_api_calls", 0) or 0,
}
comp = getattr(agent, "context_compressor", None)
if comp:
ctx_used = getattr(comp, "last_prompt_tokens", 0) or usage["total"] or 0
ctx_max = getattr(comp, "context_length", 0) or 0
if ctx_max:
usage["context_used"] = ctx_used
usage["context_max"] = ctx_max
usage["context_percent"] = max(0, min(100, round(ctx_used / ctx_max * 100)))
usage = _context_usage_for_tool_events(agent, messages)
for tc, name, args, middleware_trace, block_result, blocked_by_guardrail in parsed_calls:
if block_result is not None:
continue
Expand Down Expand Up @@ -732,21 +747,7 @@ def _run_tool(index, tool_call, function_name, function_args, middleware_trace):
agent._touch_activity(f"tool completed: {name} ({tool_duration:.1f}s)")

# Inject context usage into tool.complete events for real-time updates.
_comp = getattr(agent, "context_compressor", None)
_usage = {
"model": getattr(agent, "model", ""),
"input": getattr(agent, "session_input_tokens", 0) or 0,
"output": getattr(agent, "session_output_tokens", 0) or 0,
"total": getattr(agent, "session_total_tokens", 0) or 0,
"calls": getattr(agent, "session_api_calls", 0) or 0,
}
if _comp:
_ctx_used = getattr(_comp, "last_prompt_tokens", 0) or _usage["total"] or 0
_ctx_max = getattr(_comp, "context_length", 0) or 0
if _ctx_max:
_usage["context_used"] = _ctx_used
_usage["context_max"] = _ctx_max
_usage["context_percent"] = max(0, min(100, round(_ctx_used / _ctx_max * 100)))
_usage = _context_usage_for_tool_events(agent, messages)

if not blocked and agent.tool_complete_callback:
try:
Expand Down Expand Up @@ -932,21 +933,7 @@ def execute_tool_calls_sequential(agent, assistant_message, messages: list, effe

if not _execution_blocked and agent.tool_start_callback:
# Inject context usage for real-time context bar updates.
_t_usage = {
"model": getattr(agent, "model", ""),
"input": getattr(agent, "session_input_tokens", 0) or 0,
"output": getattr(agent, "session_output_tokens", 0) or 0,
"total": getattr(agent, "session_total_tokens", 0) or 0,
"calls": getattr(agent, "session_api_calls", 0) or 0,
}
_t_comp = getattr(agent, "context_compressor", None)
if _t_comp:
_t_ctx_used = getattr(_t_comp, "last_prompt_tokens", 0) or _t_usage["total"] or 0
_t_ctx_max = getattr(_t_comp, "context_length", 0) or 0
if _t_ctx_max:
_t_usage["context_used"] = _t_ctx_used
_t_usage["context_max"] = _t_ctx_max
_t_usage["context_percent"] = max(0, min(100, round(_t_ctx_used / _t_ctx_max * 100)))
_t_usage = _context_usage_for_tool_events(agent, messages)
try:
agent.tool_start_callback(tool_call.id, function_name, function_args, usage=_t_usage)
except Exception as cb_err:
Expand Down Expand Up @@ -1385,21 +1372,7 @@ def _execute(next_args: dict) -> Any:

if not _execution_blocked and agent.tool_complete_callback:
# Inject context usage for real-time context bar updates.
_sc_usage = {
"model": getattr(agent, "model", ""),
"input": getattr(agent, "session_input_tokens", 0) or 0,
"output": getattr(agent, "session_output_tokens", 0) or 0,
"total": getattr(agent, "session_total_tokens", 0) or 0,
"calls": getattr(agent, "session_api_calls", 0) or 0,
}
_sc_comp = getattr(agent, "context_compressor", None)
if _sc_comp:
_sc_ctx_used = getattr(_sc_comp, "last_prompt_tokens", 0) or _sc_usage["total"] or 0
_sc_ctx_max = getattr(_sc_comp, "context_length", 0) or 0
if _sc_ctx_max:
_sc_usage["context_used"] = _sc_ctx_used
_sc_usage["context_max"] = _sc_ctx_max
_sc_usage["context_percent"] = max(0, min(100, round(_sc_ctx_used / _sc_ctx_max * 100)))
_sc_usage = _context_usage_for_tool_events(agent, messages)
try:
agent.tool_complete_callback(tool_call.id, function_name, function_args, function_result, usage=_sc_usage)
except Exception as cb_err:
Expand Down
65 changes: 64 additions & 1 deletion apps/desktop/src/app/session/hooks/use-message-stream.test.tsx
Original file line number Diff line number Diff line change
Expand Up @@ -5,7 +5,7 @@ import { afterEach, beforeEach, describe, expect, it, vi } from 'vitest'

import type { ClientSessionState } from '@/app/types'
import { createClientSessionState } from '@/lib/chat-runtime'
import { $currentUsage } from '@/store/session'
import { $currentUsage, $sessionActivityStatus } from '@/store/session'
import type { RpcEvent } from '@/types/hermes'

import { useMessageStream } from './use-message-stream'
Expand Down Expand Up @@ -107,3 +107,66 @@ describe('useMessageStream token usage events', () => {
expect($currentUsage.get()).toEqual({ calls: 1, input: 10, output: 5, total: 15 })
})
})

describe('useMessageStream status.update events', () => {
beforeEach(() => {
handleEvent = () => undefined
$sessionActivityStatus.set(null)
})

afterEach(() => {
cleanup()
$sessionActivityStatus.set(null)
vi.restoreAllMocks()
})

const statusEvent = (text: string, kind = 'lifecycle', sessionId = 'session-1') =>
({
payload: { kind, text },
session_id: sessionId,
type: 'status.update'
}) as RpcEvent

it('surfaces lifecycle statuses for the active session', () => {
render(<MessageStreamHarness />)

act(() => handleEvent(statusEvent('πŸ“¦ Preflight compression: ~90,000 tokens. This may take a moment.')))

expect($sessionActivityStatus.get()).toEqual({
kind: 'lifecycle',
text: 'πŸ“¦ Preflight compression: ~90,000 tokens. This may take a moment.'
})
})

it('ignores statuses from inactive sessions and unknown kinds', () => {
render(<MessageStreamHarness />)

act(() => handleEvent(statusEvent('background noise', 'lifecycle', 'session-2')))
expect($sessionActivityStatus.get()).toBeNull()

act(() => handleEvent(statusEvent('voice things', 'voice')))
expect($sessionActivityStatus.get()).toBeNull()
})

it('clears on ready kind and on stream activity', () => {
render(<MessageStreamHarness />)

act(() => handleEvent(statusEvent('β ‹ compressing 120 messages', 'compressing')))
expect($sessionActivityStatus.get()).not.toBeNull()

act(() => handleEvent(statusEvent('ready', 'ready')))
expect($sessionActivityStatus.get()).toBeNull()

act(() => handleEvent(statusEvent('πŸ“¦ Compression complete: ~90,000 β†’ ~30,000 tokens.')))
expect($sessionActivityStatus.get()).not.toBeNull()

act(() =>
handleEvent({
payload: { text: 'hello' },
session_id: 'session-1',
type: 'message.delta'
} as RpcEvent)
)
expect($sessionActivityStatus.get()).toBeNull()
})
})
Loading
Loading