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
21 changes: 21 additions & 0 deletions archytas/agent.py
Original file line number Diff line number Diff line change
Expand Up @@ -441,6 +441,27 @@ async def build_tail_injections(self) -> list[BaseMessage]:
fired.append((tool_name, output))

if fired:
# Record the last-known token size of each fired statetool's output
# (issue #86). These injections are ephemeral and never persisted,
# so ChatHistory tracks them separately (statetool_token_estimates)
# and folds them into token_overhead. Entries are only ever
# updated here — a statetool that doesn't fire this turn keeps its
# previous size rather than dropping to zero — which keeps token
# estimates stable across loop steps and prevents the actual-usage
# delta from being wrongly absorbed into base_tokens.
for tool_name, output in fired:
try:
estimate = await self.model.get_num_tokens_from_messages(
[HumanMessage(content=output)]
)
except Exception as err:
logger.debug(
"Token estimation for statetool %r failed: %r", tool_name, err
)
continue
if estimate:
self.chat_history.statetool_token_estimates[tool_name] = estimate

state_messages = await self.build_state_injection(fired)
tail.extend(state_messages)
if self.verbose:
Expand Down
23 changes: 23 additions & 0 deletions archytas/chat_history.py
Original file line number Diff line number Diff line change
Expand Up @@ -415,6 +415,9 @@ class OutboundChatHistory:
summary_token_count: Optional[int]
overhead_token_count: Optional[int]
summarization_threshold: Optional[int]
# Last-known per-statetool injection sizes (issue #86). Optional with a
# default so existing constructors that don't pass it keep working.
statetool_token_estimates: Optional[dict[str, int]] = None


class ChatHistory:
Expand All @@ -430,6 +433,13 @@ class ChatHistory:
history_summarizer: Optional[callable]
summarization_threshold: int
tool_token_estimate: int
# Last-known token size of each statetool's injected output, keyed by tool
# name (issue #86). Statetool pairs are ephemeral tail injections, so they
# never appear in the persisted records; tracking their last observed size
# here keeps the token overhead stable across loop steps where a statetool
# does or doesn't fire. Entries are updated when a statetool fires and
# retained (not zeroed) when it doesn't.
statetool_token_estimates: dict[str, int]
_tool_hash: str
_token_estimate: int|None
history_summarization_task: Optional[asyncio.Task]
Expand Down Expand Up @@ -459,6 +469,7 @@ def __init__(
self.system_message = None
self.model = model
self.tool_token_estimate = 0
self.statetool_token_estimates = {}
self._tool_hash = ""
self._token_estimate = None
self.loop_summarizer = loop_summarizer
Expand Down Expand Up @@ -612,11 +623,20 @@ def instruction_token_estimate(self) -> int:
return self.instruction.token_count
return 0

@property
def statetool_token_estimate(self) -> int:
"""Sum of the last-known token sizes of all statetool injections."""
return sum(
value for value in self.statetool_token_estimates.values()
if isinstance(value, int)
)

@property
def token_overhead(self):
return sum((value for value in (
self.base_tokens,
self.tool_token_estimate,
self.statetool_token_estimate,
self.instruction_token_estimate,
) if isinstance(value, int)))

Expand Down Expand Up @@ -928,6 +948,7 @@ def to_dict(self) -> dict[str, Any]:
"current_loop_id": self.current_loop_id,
"summarization_threshold": self.summarization_threshold,
"tool_token_estimate": self.tool_token_estimate,
"statetool_token_estimates": dict(self.statetool_token_estimates),
"token_estimate": self._token_estimate,
},
"system_message": self.system_message.to_dict() if self.system_message else None,
Expand Down Expand Up @@ -983,6 +1004,8 @@ def from_dict(
history.summarization_threshold = metadata["summarization_threshold"]
if metadata.get("tool_token_estimate") is not None:
history.tool_token_estimate = metadata["tool_token_estimate"]
if metadata.get("statetool_token_estimates"):
history.statetool_token_estimates = dict(metadata["statetool_token_estimates"])
if metadata.get("token_estimate") is not None:
history._token_estimate = metadata["token_estimate"]

Expand Down
242 changes: 230 additions & 12 deletions archytas/summarizers.py
Original file line number Diff line number Diff line change
@@ -1,3 +1,25 @@
"""
Summarization pipeline for context management.

Two independent gates shrink the outgoing chat history, and keeping them
separate is what makes the flow tractable:

1. Selection (threshold): ``get_records_up_to_threshold`` decides *which*
records are eligible for summarization, based on ``token_threshold``. It
never truncates content — it only chooses where the summarize/keep boundary
sits.
2. Fitting (context window): ``fit_records_to_context`` runs on the records the
selection step chose, just before the summarization request, and
middle-truncates copies of any records that, together, overflow the model's
context window. Records that fit are passed through untouched, and the
original records are never mutated.

The distinction matters: a record larger than the summarization threshold but
smaller than the context window is summarized in full and never truncated.
Truncation happens only when the *selected* records don't fit the context
window (see ``fit_records_to_context`` for the exact budget).
"""

import asyncio
import json
import logging
Expand Down Expand Up @@ -27,6 +49,23 @@
MESSAGE_SUMMARIZATION_THRESHOLD: int = 1000
MESSAGE_SUMMARIZATION_SNIPPET_SIZE: int = 1000

# Headroom reserved out of the model's context window when sizing a
# summarization request: covers the system/user template scaffolding plus the
# generated summary output.
SUMMARIZATION_TOKEN_RESERVE: int = 8192
# Below this per-record token target, middle-truncation would leave too little
# content to be meaningful; fall back to a placeholder instead.
MIN_TRUNCATION_TOKENS: int = 256

TRUNCATION_NOTICE_TEMPLATE = (
"\n\n[... {removed} characters removed from the middle of this message "
"because it was too large to summarize in full ...]\n\n"
)
TRUNCATION_PLACEHOLDER = (
"[Message content omitted from summarization because it was too large to "
"fit in the context window.]"
)


async def get_summarizable_records(chat_history: "ChatHistory") -> "list[RecordType]":
from .chat_history import RecordType, SummaryRecord, SystemMessage, AIMessage
Expand All @@ -44,33 +83,206 @@ async def get_summarizable_records(chat_history: "ChatHistory") -> "list[RecordT
return summarizable_records


async def get_records_up_to_threshold(chat_history: "ChatHistory", model: BaseArchytasModel, token_threshold: int) -> "list[RecordType]":
def _has_pending_tool_calls(message: "BaseMessage") -> bool:
"""True if `message` is an AIMessage that issued tool calls (whose
ToolMessage responses therefore appear in later records)."""
from .chat_history import AIMessage
return isinstance(message, AIMessage) and bool(message.tool_calls)


def retreat_past_pending_tool_calls(records: "list[RecordType]", idx: int) -> int:
"""Move `idx` backward so that ``records[:idx+1]`` doesn't end on an
AIMessage whose ToolMessage responses would be left out. Returns the new
boundary index (may be -1 if no complete boundary exists at or before
`idx`)."""
while idx >= 0 and _has_pending_tool_calls(records[idx].message):
idx -= 1
return idx


def extend_to_include_tool_responses(records: "list[RecordType]", idx: int) -> int:
"""Move `idx` forward so that ``records[:idx+1]`` includes the ToolMessage
responses following an AIMessage's tool calls, avoiding a split between an
AIMessage and its responses. Returns the new boundary index."""
from .chat_history import ToolMessage
while idx < len(records) - 1 and (
_has_pending_tool_calls(records[idx].message)
or isinstance(records[idx + 1].message, ToolMessage)
):
idx += 1
return idx


async def get_records_up_to_threshold(chat_history: "ChatHistory", model: BaseArchytasModel, token_threshold: int) -> "list[RecordType]":
"""Select the prefix of summarizable records whose cumulative token count
stays under ``token_threshold``.

Note: the returned prefix may extend *past* the threshold. When the very
first summarizable record is already larger than the whole threshold, the
normal walk would select nothing and summarization could never make
progress, so the fallback branch selects the minimal oversized prefix
instead. This function never truncates content — fitting the selected
records into the context window (including any truncation) is handled
separately by ``fit_records_to_context``.
"""
token_count = chat_history.base_tokens or 0
target_idx: int = -1
threshold_exceeded = False
summarizable_records = await get_summarizable_records(chat_history)
for idx, record in enumerate(summarizable_records):
token_count += record.token_count
token_count += record.token_count or 0
if token_count >= token_threshold:
threshold_exceeded = True
break
target_idx = idx
else:
# The whole history fits under the threshold; nothing to summarize.
target_idx = -1

# Don't split AIMessages and their tool call
done = False
while target_idx >= 0 and not done:
target_record = summarizable_records[target_idx]
target_message = target_record.message
if isinstance(target_message, AIMessage) and target_message.tool_calls:
target_idx -= 1
else:
done = True
# Don't end the selection on an AIMessage whose tool responses would be split off.
target_idx = retreat_past_pending_tool_calls(summarizable_records, target_idx)

if target_idx >= 0:
return summarizable_records[:target_idx+1]

# Fallback (issue #85): a single record at the head of the history is
# larger than the whole summarization threshold, so the walk above selected
# nothing. Without this, summarization triggers on every query but never
# makes progress. Select the minimal prefix instead — the head record plus
# however many records are needed to avoid splitting an AIMessage from its
# ToolMessage responses — and let the summarizer shrink any content that
# doesn't fit in the context window (see fit_records_to_context).
if threshold_exceeded and summarizable_records:
end_idx = extend_to_include_tool_responses(summarizable_records, 0)
selected = summarizable_records[:end_idx + 1]
logger.warning(
"A single message (%s tokens) exceeds the summarization threshold "
"(%s tokens); summarizing %d record(s) beyond the threshold, "
"truncating content as needed to fit the context window.",
selected[0].token_count, token_threshold, len(selected),
)
return selected

return []


def middle_truncate_text(text: str, keep_chars: int) -> str:
"""Cut characters out of the middle of `text` so roughly `keep_chars`
characters remain, keeping the head and tail and inserting a notice where
content was removed."""
if keep_chars <= 0:
return TRUNCATION_PLACEHOLDER
if len(text) <= keep_chars:
return text
head = keep_chars // 2
tail = keep_chars - head
notice = TRUNCATION_NOTICE_TEMPLATE.format(removed=len(text) - keep_chars)
return text[:head] + notice + text[-tail:]


async def shrink_record_for_summarization(
record: "MessageRecord",
model: BaseArchytasModel,
target_tokens: int,
) -> "MessageRecord":
"""
Return a copy of `record` whose message content has been middle-truncated
to fit within `target_tokens` (issue #85 remediation b, with placeholder
replacement as the fallback). The copy keeps the original record's uuid so
the resulting summary excludes the original — which is never mutated — from
future outgoing history.
"""
from langchain_core.messages import HumanMessage
from .chat_history import MessageRecord

message = record.message.model_copy(deep=True)
content = message.content
new_content: str
token_count: int

if isinstance(content, str) and target_tokens >= MIN_TRUNCATION_TOKENS:
token_count = record.token_count or 0
new_content = content
# Estimate a keep-size from the observed chars-per-token ratio, then
# re-estimate and tighten a few times if the cut wasn't deep enough.
for _ in range(4):
chars_per_token = len(new_content) / max(token_count, 1)
keep_chars = int(target_tokens * chars_per_token * 0.9)
new_content = middle_truncate_text(content, keep_chars)
token_count = await model.get_num_tokens_from_messages(
[HumanMessage(content=new_content)]
)
if token_count <= target_tokens:
break
else:
new_content = TRUNCATION_PLACEHOLDER
token_count = await model.get_num_tokens_from_messages(
[HumanMessage(content=new_content)]
)
else:
return []
# Non-string content (multimodal/structured blocks) has no
# well-defined middle to cut; replace it outright.
new_content = TRUNCATION_PLACEHOLDER
token_count = await model.get_num_tokens_from_messages(
[HumanMessage(content=new_content)]
)

message.content = new_content
return MessageRecord(
message=message,
uuid=record.uuid,
token_count=token_count,
metadata=dict(record.metadata),
react_loop_id=record.react_loop_id,
)


async def fit_records_to_context(
records: "list[MessageRecord]",
model: BaseArchytasModel,
) -> "list[MessageRecord]":
"""
Ensure the records destined for a summarization request fit within the
model's context window (minus SUMMARIZATION_TOKEN_RESERVE headroom).
Oversized records are replaced by middle-truncated copies (originals are
left untouched); records that fit are passed through unchanged.
"""
context_window = model.contextsize()
if not context_window or not records:
return records

# Reserve headroom for the summarization prompt scaffolding + generated
# summary, but never let that reserve claim more than half the window: on a
# small context window SUMMARIZATION_TOKEN_RESERVE could otherwise dominate
# (or drive the budget negative), leaving no room to summarize at all.
budget = max(context_window - SUMMARIZATION_TOKEN_RESERVE, context_window // 2)
fitted = list(records)
total = sum(record.token_count or 0 for record in fitted)

attempts = 0
while total > budget and attempts < len(fitted) * 2:
attempts += 1
idx, largest = max(
enumerate(fitted), key=lambda pair: pair[1].token_count or 0
)
largest_tokens = largest.token_count or 0
if largest_tokens <= 0:
break
overage = total - budget
target_tokens = max(largest_tokens - overage, 0)
logger.warning(
"Record %s (%s tokens) is too large to summarize in full; "
"truncating its copy to ~%s tokens for the summarization request.",
largest.uuid, largest_tokens, target_tokens,
)
shrunk = await shrink_record_for_summarization(largest, model, target_tokens)
if (shrunk.token_count or 0) >= largest_tokens:
# No progress is possible on this record; give up rather than loop.
break
fitted[idx] = shrunk
total = sum(record.token_count or 0 for record in fitted)

return fitted


async def get_records_up_to_loop(chat_history: "ChatHistory", model: BaseArchytasModel) -> "list[RecordType]":
Expand Down Expand Up @@ -132,6 +344,12 @@ async def default_history_summarizer(
case MessageRecord():
records_to_summarize.append(record)

# Shrink any content too large to fit the summarization request into the
# context window (issue #85). Only the copies sent to the summarizer are
# truncated; the original records are untouched and are excluded from
# future outgoing history once the summary lands.
records_to_summarize = await fit_records_to_context(records_to_summarize, model)

uuids = [record.uuid for record in records_to_summarize]

jinja_globals = {
Expand Down
Loading
Loading