Skip to content
Open
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
72 changes: 53 additions & 19 deletions astrbot/core/agent/context/config.py
Original file line number Diff line number Diff line change
@@ -1,5 +1,5 @@
from dataclasses import dataclass
from typing import TYPE_CHECKING
from typing import TYPE_CHECKING, Literal

from .compressor import ContextCompressor
from .token_counter import TokenCounter
Expand All @@ -10,25 +10,59 @@

@dataclass
class ContextConfig:
"""Context configuration class."""

max_context_tokens: int = 0
"""Maximum number of context tokens. <= 0 means no limit."""
enforce_max_turns: int = -1 # -1 means no limit
"""Maximum number of conversation turns to keep. -1 means no limit. Executed before compression."""
truncate_turns: int = 1
"""Number of conversation turns to discard at once when truncation is triggered.
Two processes will use this value:

1. Enforce max turns truncation.
2. Truncation by turns compression strategy.
"""Context configuration class — orthogonal trigger/disposal model.

Trigger dimension (WHEN) — checked independently:
enable_turn_limit / max_turns
enable_token_guard / token_guard_threshold

Disposal dimension (WHAT) — executed in order when any trigger fires:
1. summary (if enabled and provider available)
2. discard (fallback if summary fails or is disabled)

Retention constraint — lower bound on how many turns remain after discard:
retain_turns / retain_percentage. Discard will not remove turns below
this floor.
Double-check halving — unconditional truncation when still over token
threshold after disposal (only fired when enable_token_guard is True).
"""
llm_compress_instruction: str | None = None
"""Instruction prompt for LLM-based compression."""
llm_compress_keep_recent_ratio: float = 0.15
"""Percent of current context tokens to keep as exact recent context during LLM-based compression."""
llm_compress_provider: "Provider | None" = None
"""LLM provider used for compression tasks. If None, truncation strategy is used."""

# -- Trigger dimension --
enable_turn_limit: bool = False
"""Enable turn-based trigger. When True, exceeding max_turns triggers disposal."""
max_turns: int = 50
"""Maximum conversation turns before disposal is triggered. Must be >= 2."""
enable_token_guard: bool = True
"""Enable token-count trigger. When True, exceeding token_guard_threshold
triggers disposal."""
token_guard_threshold: float = 0.82
"""Token usage ratio (current_tokens / max_tokens) that triggers disposal.
Range 0.5–0.99."""

# -- Disposal dimension (compression behavior) --
enable_summary: bool = True
"""Enable LLM-based summary compression. Takes priority over discard when
both are enabled."""
enable_discard: bool = True
"""Enable discard of oldest turns. Used as fallback if summary fails or
is disabled."""
discard_turns: int = 1
"""Number of turns to discard at once. Must be >= 1."""
summary_prompt: str = ""
"""Custom instruction prompt for summary generation. Empty = use built-in."""
summary_provider: "Provider | None" = None
"""Resolved LLM provider for summary generation. None = no summary available."""

# -- Retention (lower bound) --
retention_method: Literal["turns", "percentage", "null"] = "turns"
"""Retention method: 'turns', 'percentage', or 'null'."""
retain_turns: int = 20
"""Minimum turns to keep when retention_method is 'turns'. Must be >= 1."""
retain_percentage: float = 0.3
"""Minimum ratio of turns to keep when retention_method is 'percentage'.
Range 0.1–0.9."""

# -- Customisation --
custom_token_counter: TokenCounter | None = None
"""Custom token counting method. If None, the default method is used."""
custom_compressor: ContextCompressor | None = None
Expand Down
256 changes: 175 additions & 81 deletions astrbot/core/agent/context/manager.py
Original file line number Diff line number Diff line change
@@ -1,121 +1,215 @@
import asyncio
import math

from astrbot import logger

from ..message import Message
from .compressor import LLMSummaryCompressor, TruncateByTurnsCompressor
from .compressor import LLMSummaryCompressor
from .config import ContextConfig
from .round_utils import split_into_rounds
from .token_counter import EstimateTokenCounter
from .truncator import ContextTruncator


class ContextManager:
"""Context compression manager."""
"""Context compression manager — orthogonal trigger/disposal model."""

def __init__(
self,
config: ContextConfig,
) -> None:
"""Initialize the context manager.

There are two strategies to handle context limit reached:
1. Truncate by turns: remove older messages by turns.
2. LLM-based compression: use LLM to summarize old messages.

Args:
config: The context configuration.
"""
self.config = config

self.token_counter = config.custom_token_counter or EstimateTokenCounter()
self.truncator = ContextTruncator()

# Build compressors on demand. Summary compressor only when a provider is
# available; discard is handled directly via truncator, not a compressor.
self._summary_compressor = None
self._unity_compressor = None
if config.custom_compressor:
self.compressor = config.custom_compressor
elif config.llm_compress_provider:
self.compressor = LLMSummaryCompressor(
provider=config.llm_compress_provider,
keep_recent_ratio=config.llm_compress_keep_recent_ratio,
instruction_text=config.llm_compress_instruction,
token_counter=self.token_counter,
)
self._unity_compressor = config.custom_compressor
else:
self.compressor = TruncateByTurnsCompressor(
truncate_turns=config.truncate_turns
if config.summary_provider:
self._summary_compressor = LLMSummaryCompressor(
provider=config.summary_provider,
keep_recent_ratio=config.retain_percentage,
instruction_text=config.summary_prompt,
compression_threshold=config.token_guard_threshold,
token_counter=self.token_counter,
)

# -- helpers ----------------------------------------------------------

def _count_turns(self, messages: list[Message]) -> int:
"""Count the number of conversation turns (non-system rounds)."""
rounds = split_into_rounds(messages)
# Filter out rounds that are exclusively system messages
Comment thread
sourcery-ai[bot] marked this conversation as resolved.
return sum(
1
for rnd in rounds
if any((isinstance(m, Message) and m.role != "system") for m in rnd)
)

def _compute_discard_limit(self, total_turns: int) -> int:
"""Maximum number of turns that discard may remove, given retention."""
method = self.config.retention_method
if method == "turns":
return max(0, total_turns - self.config.retain_turns)
if method == "percentage":
min_keep = math.ceil(total_turns * self.config.retain_percentage)
return max(0, total_turns - min_keep)
# "null" — no lower bound
return total_turns

async def _try_summary(self, messages: list[Message]) -> list[Message] | None:
"""Attempt LLM summary compression. Returns compressed messages or None."""
if not self.config.enable_summary or self._summary_compressor is None:
return None
try:
result = await self._summary_compressor(messages)
if result is None or result is messages:
return None # compressor chose not to compress
if len(result) >= len(messages):
return None # no effective reduction
return result
except asyncio.CancelledError:
raise
except Exception:
logger.warning(
Comment thread
sourcery-ai[bot] marked this conversation as resolved.
"LLM summary compression failed, falling back.", exc_info=True
)
return None

def _try_discard(
self, messages: list[Message], total_turns: int
) -> list[Message] | None:
"""Discard oldest turns, bounded by retention."""
if not self.config.enable_discard:
return None
max_discardable = self._compute_discard_limit(total_turns)
if max_discardable <= 0:
return None # retention prevents any discard
requested = min(self.config.discard_turns, max_discardable)
return self.truncator.truncate_by_dropping_oldest_turns(
messages,
drop_turns=requested,
)

async def process(
self, messages: list[Message], trusted_token_usage: int = 0
) -> list[Message]:
"""Process the messages.
def _token_guard_exceeded(self, tokens: int, max_context_tokens: int) -> bool:
Comment thread
sourcery-ai[bot] marked this conversation as resolved.
"""Pure predicate: return True if tokens exceed the configured threshold."""
if tokens <= 0 or max_context_tokens <= 0:
return False
return (tokens / max_context_tokens) > self.config.token_guard_threshold

Args:
messages: The original message list.
def _triggers_fired(
self,
total_turns: int,
current_tokens: int,
guard_max_tokens: int,
) -> bool:
"""Return True if any trigger condition is met.

Returns:
The processed message list.
Args:
total_turns: Current conversation turn count.
current_tokens: Current token count.
guard_max_tokens: Resolved guard window (0 = token guard disabled).
"""
try:
result = messages
if self.config.enable_turn_limit and total_turns > self.config.max_turns:
return True

# 1. 基于轮次的截断 (Enforce max turns)
if self.config.enforce_max_turns != -1:
result = self.truncator.truncate_by_turns(
result,
keep_most_recent_turns=self.config.enforce_max_turns,
drop_turns=self.config.truncate_turns,
)
if (
self.config.enable_token_guard
and guard_max_tokens
and self._token_guard_exceeded(current_tokens, guard_max_tokens)
):
return True

# 2. 基于 token 的压缩
if self.config.max_context_tokens > 0:
total_tokens = self.token_counter.count_tokens(
result, trusted_token_usage
)
return False

if self.compressor.should_compress(
result, total_tokens, self.config.max_context_tokens
):
result = await self._run_compression(result, total_tokens)
async def _select_disposal(
self,
messages: list[Message],
total_turns: int,
) -> list[Message]:
"""Apply disposal strategy: custom > summary > discard."""
if self._unity_compressor is not None:
return await self._unity_compressor(messages)

return result
except Exception as e:
logger.error(f"Error during context processing: {e}", exc_info=True)
return messages
compressed = await self._try_summary(messages)
if compressed is not None:
return compressed

discarded = self._try_discard(messages, total_turns=total_turns)
if discarded is not None:
return discarded

async def _run_compression(
self, messages: list[Message], prev_tokens: int
logger.warning(
"Context disposal triggered but both summary and discard "
"are unavailable or disabled. No compression applied.",
)
return messages

# -- main entry point -------------------------------------------------

async def process(
self,
messages: list[Message],
trusted_token_usage: int = 0,
max_context_tokens: int = 0,
) -> list[Message]:
"""
Compress/truncate the messages.
"""Process messages through the orthogonal trigger/disposal pipeline.

Args:
messages: The original message list.
prev_tokens: The token count before compression.
trusted_token_usage: External token count hint (e.g. from conversation stats).
max_context_tokens: The model's context window size (provider-level value).

Returns:
The compressed/truncated message list.
The processed message list.
"""
logger.debug("Compress triggered, starting compression...")

messages = await self.compressor(messages)

# double check
tokens_after_summary = self.token_counter.count_tokens(messages)

# calculate compress rate
compress_rate = (tokens_after_summary / self.config.max_context_tokens) * 100
logger.info(
f"Compress completed."
f" {prev_tokens} -> {tokens_after_summary} tokens,"
f" compression rate: {compress_rate:.2f}%.",
)
try:
result = messages

# last check
if self.compressor.should_compress(
messages, tokens_after_summary, self.config.max_context_tokens
):
logger.info(
"Context still exceeds max tokens after compression, applying halving truncation..."
current_tokens = self.token_counter.count_tokens(
result, trusted_token_usage
)
# still need compress, truncate by half
messages = self.truncator.truncate_by_halving(messages)
total_turns = self._count_turns(result)

# Resolve token guard: validate config and precompute guard window.
guard_max_tokens = 0
if self.config.enable_token_guard:
if max_context_tokens <= 0:
if not getattr(self, "_token_guard_warning_emitted", False):
logger.warning(
"Token guard is enabled but max_context_tokens is %s. "
"Token guarding is effectively disabled. "
"Set max_context_tokens in the provider config to enable it.",
max_context_tokens,
)
self._token_guard_warning_emitted = True
else:
guard_max_tokens = max_context_tokens

if not self._triggers_fired(total_turns, current_tokens, guard_max_tokens):
return result

# Disposal entry: custom > summary > discard (fallthrough)
result = await self._select_disposal(result, total_turns)

# Double-check halving — only when token guard is active
tokens_after = self.token_counter.count_tokens(result, trusted_token_usage)
if guard_max_tokens and self._token_guard_exceeded(
tokens_after, guard_max_tokens
):
logger.info(
"Context still exceeds token guard threshold after disposal, "
"applying halving truncation (unconstrained by retention).",
)
result = self.truncator.truncate_by_halving(result)

return messages
return result
except asyncio.CancelledError:
raise
except Exception:
logger.error("Error during context processing.", exc_info=True)
return messages
Loading
Loading