Skip to content
Draft
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
57 changes: 53 additions & 4 deletions astrbot/core/agent/context/compressor.py
Original file line number Diff line number Diff line change
@@ -1,10 +1,14 @@
import asyncio
from collections.abc import Awaitable, Callable
from typing import TYPE_CHECKING, Protocol, runtime_checkable

from ...provider.modalities import (
log_context_sanitize_stats,
sanitize_contexts_by_modalities,
)
from ..message import Message
from ..event_stream import RequestEventRecorder, request_recorder_kwargs
from ..message import Message, dump_messages_with_checkpoints
from ..response import AgentResponse
from .token_counter import EstimateTokenCounter, TokenCounter

if TYPE_CHECKING:
Expand Down Expand Up @@ -130,6 +134,8 @@ def __init__(
instruction_text: str | None = None,
compression_threshold: float = 0.82,
token_counter: TokenCounter | None = None,
request_event_emitter: Callable[[AgentResponse], Awaitable[None]] | None = None,
turn_id: str | None = None,
) -> None:
"""Initialize the LLM summary compressor.

Expand All @@ -139,8 +145,13 @@ def __init__(
exact context. Clamped to 0-0.3.
instruction_text: Custom instruction for summary generation.
compression_threshold: The compression trigger threshold (default: 0.82).
token_counter: Optional context token estimator.
request_event_emitter: Optional acknowledged runner event sink.
turn_id: Runner turn identity attached to request events.
"""
self.provider = provider
self.request_event_emitter = request_event_emitter
self.turn_id = turn_id
self.keep_recent_ratio = min(max(float(keep_recent_ratio), 0.0), 0.3)
self.compression_threshold = compression_threshold
self.token_counter = token_counter or EstimateTokenCounter()
Expand Down Expand Up @@ -247,6 +258,16 @@ async def __call__(self, messages: list[Message]) -> list[Message]:
if not any(msg.role != "system" for msg in summary_contexts):
return messages

summary_contexts = [
Message.model_validate(item)
for item in dump_messages_with_checkpoints(
[message for message in summary_contexts if not message._no_save]
)
if item.get("role") != "_checkpoint"
]
if not summary_contexts:
return messages

if summary_contexts[-1].role != "assistant":
summary_contexts.append(
Message(
Expand All @@ -273,9 +294,37 @@ async def __call__(self, messages: list[Message]) -> list[Message]:

# Generate summary
try:
response = await self.provider.text_chat(
contexts=sanitized_summary_contexts,
)
recorder = None
if self.request_event_emitter:
recorder = RequestEventRecorder(
self.request_event_emitter,
{
"turn_id": self.turn_id,
"purpose": "compaction",
"provider_id": self.provider.provider_config.get("id", ""),
"model": self.provider.get_model(),
},
)
await recorder.begin()
try:
response = await self.provider.text_chat(
contexts=sanitized_summary_contexts,
**request_recorder_kwargs(self.provider.text_chat, recorder),
)
if recorder:
await recorder.finish(
"failed" if response.role == "err" else "completed",
usage=response.usage,
)
except BaseException as exc:
if recorder:
await recorder.finish(
"cancelled"
if isinstance(exc, asyncio.CancelledError)
else "failed",
error_code=type(exc).__name__,
)
raise
summary_content = (response.completion_text or "").strip()
except Exception as e:
logger.error(f"Failed to generate summary: {e}")
Expand Down
6 changes: 6 additions & 0 deletions astrbot/core/agent/context/config.py
Original file line number Diff line number Diff line change
@@ -1,6 +1,8 @@
from collections.abc import Awaitable, Callable
from dataclasses import dataclass
from typing import TYPE_CHECKING

from ..response import AgentResponse
from .compressor import ContextCompressor
from .token_counter import TokenCounter

Expand Down Expand Up @@ -33,3 +35,7 @@ class ContextConfig:
"""Custom token counting method. If None, the default method is used."""
custom_compressor: ContextCompressor | None = None
"""Custom context compression method. If None, the default method is used."""
request_event_emitter: Callable[[AgentResponse], Awaitable[None]] | None = None
"""Optional runtime event sink for summary model requests."""
turn_id: str | None = None
"""Runner turn identity for summary request events."""
2 changes: 2 additions & 0 deletions astrbot/core/agent/context/manager.py
Original file line number Diff line number Diff line change
Expand Up @@ -36,6 +36,8 @@ def __init__(
keep_recent_ratio=config.llm_compress_keep_recent_ratio,
instruction_text=config.llm_compress_instruction,
token_counter=self.token_counter,
request_event_emitter=config.request_event_emitter,
turn_id=config.turn_id,
)
else:
self.compressor = TruncateByTurnsCompressor(
Expand Down
Loading