From f54063abd5dd8b37b0a00418e281c30458476e04 Mon Sep 17 00:00:00 2001 From: C10H14N2O5 <100066858+C10H14N2O5@users.noreply.github.com> Date: Mon, 24 Aug 2026 17:19:09 +0800 Subject: [PATCH 01/16] feat: add manual context compression command --- .../builtin_commands/commands/conversation.py | 230 +++++++- .../builtin_stars/builtin_commands/main.py | 5 + astrbot/core/agent/context/compressor.py | 46 +- astrbot/core/agent/context/config.py | 2 + astrbot/core/agent/context/manager.py | 57 +- astrbot/core/astr_main_agent.py | 23 +- tests/agent/test_context_manager.py | 144 ++++- tests/test_conversation_commands.py | 538 ++++++++++++++++++ tests/unit/test_astr_main_agent.py | 83 +++ 9 files changed, 1097 insertions(+), 31 deletions(-) diff --git a/astrbot/builtin_stars/builtin_commands/commands/conversation.py b/astrbot/builtin_stars/builtin_commands/commands/conversation.py index 3d3edb911b..3443f64bc2 100644 --- a/astrbot/builtin_stars/builtin_commands/commands/conversation.py +++ b/astrbot/builtin_stars/builtin_commands/commands/conversation.py @@ -1,16 +1,29 @@ +import json + from sqlalchemy import case, func, select from sqlmodel import col from astrbot.api import sp, star -from astrbot.api.event import AstrMessageEvent, MessageEventResult +from astrbot.api.event import AstrMessageEvent, MessageChain, MessageEventResult +from astrbot.api.message_components import Json from astrbot.core import logger +from astrbot.core.agent.context.config import ContextConfig +from astrbot.core.agent.context.manager import ContextManager +from astrbot.core.agent.context.round_utils import split_into_rounds +from astrbot.core.agent.message import ( + bind_checkpoint_messages, + dump_messages_with_checkpoints, +) +from astrbot.core.agent.response import AgentStats from astrbot.core.agent.runners.deerflow.constants import ( DEERFLOW_PROVIDER_TYPE, DEERFLOW_THREAD_ID_KEY, ) from astrbot.core.agent.runners.deerflow.deerflow_api_client import DeerFlowAPIClient +from astrbot.core.astr_main_agent import get_context_compression_provider from astrbot.core.db.po import ProviderStat from astrbot.core.utils.active_event_registry import active_event_registry +from astrbot.core.utils.session_lock import session_lock_manager THIRD_PARTY_AGENT_RUNNER_KEY = { "dify": "dify_conversation_id", @@ -166,6 +179,221 @@ async def stop(self, message: AstrMessageEvent) -> None: MessageEventResult().message("✅ No running tasks in the current session.") ) + async def compact(self, message: AstrMessageEvent) -> None: + """Compress the persisted history of the current local conversation.""" + + def reply(text: str) -> None: + """Set a plain-text command result. + + Args: + text: Message shown to the user. + """ + message.set_result(message.plain_result(text)) + + preserved = "❌ Context compression failed; the original context was preserved." + cancelled = "⚠️ Compression cancelled; original context was preserved." + unknown = "⚠️ Context state is unknown. Check the conversation before retrying." + umo = message.unified_msg_origin + cfg = self.context.get_config(umo=umo) + provider_settings = cfg.get("provider_settings", {}) + conversation_manager = self.context.conversation_manager + + is_unique_session = cfg.get("platform_settings", {}).get( + "unique_session", + False, + ) + is_shared_group = bool(message.get_group_id()) and not is_unique_session + if is_shared_group and message.role != "admin": + reply( + "❌ Context compression requires admin permission in a shared " + "group conversation." + ) + return + + if not provider_settings.get("enable", True): + reply("❌ AI features are disabled for this session.") + return + + if not provider_settings.get("enable_manual_context_compression", False): + reply( + "❌ Manual context compression is disabled. Enable it in Context " + "Management first." + ) + return + + if provider_settings.get("agent_runner_type", "local") != "local": + reply("❌ /compact is supported only by the local agent runner.") + return + + strategy = provider_settings.get("context_limit_reached_strategy") + if strategy != "llm_compress": + reply("❌ /compact requires the LLM context compression strategy.") + return + + initial_cid = await conversation_manager.get_curr_conversation_id(umo) + if not initial_cid: + reply("❌ You are not in a conversation. Use /new to create one.") + return + + compression_provider = await get_context_compression_provider( + strategy, + provider_settings.get("llm_compress_provider_id", ""), + self.context, + message, + ) + if not compression_provider: + reply("❌ No LLM provider is available for context compression.") + return + + await message.send(MessageChain().message("⏳ Compressing context...")) + + try: + async with session_lock_manager.acquire_lock(umo): + if message.is_stopped(): + return + if message.get_extra("agent_stop_requested"): + reply(cancelled) + return + + cid = await conversation_manager.get_curr_conversation_id(umo) + if not cid or cid != initial_cid: + reply("⚠️ The active conversation changed; no changes were saved.") + return + + conversation = await conversation_manager.get_conversation(umo, cid) + if not conversation: + reply( + "❌ The current conversation could not be loaded; the " + "original context was preserved." + ) + return + + original_history_text = conversation.history + original_history = json.loads(conversation.history) + if not isinstance(original_history, list) or not original_history: + reply("ℹ️ There is not enough conversation history to compress.") + return + + messages = bind_checkpoint_messages(original_history) + complete_rounds = sum( + any(segment.role == "user" for segment in round_) + and any(segment.role == "assistant" for segment in round_) + for round_ in split_into_rounds(messages) + ) + if complete_rounds <= 1: + reply("ℹ️ There is not enough conversation history to compress.") + return + + context_manager = ContextManager( + ContextConfig( + llm_compress_instruction=provider_settings.get( + "llm_compress_instruction" + ), + llm_compress_keep_recent_ratio=provider_settings.get( + "llm_compress_keep_recent_ratio", + 0.15, + ), + llm_compress_preserve_latest_round=True, + llm_compress_provider=compression_provider, + ) + ) + tokens_before = context_manager.token_counter.count_tokens(messages) + if tokens_before <= 0: + reply("ℹ️ There is not enough conversation history to compress.") + return + + compressed_messages = await context_manager.process( + messages, + force_compress=True, + ) + tokens_after = context_manager.token_counter.count_tokens( + compressed_messages + ) + if compressed_messages == messages or tokens_after >= tokens_before: + reply(preserved) + return + + target_history = dump_messages_with_checkpoints(compressed_messages) + latest_cid = await conversation_manager.get_curr_conversation_id(umo) + latest_conversation = await conversation_manager.get_conversation( + umo, cid + ) + if ( + latest_cid != cid + or not latest_conversation + or latest_conversation.history != original_history_text + ): + reply( + "⚠️ Context changed during compression; no changes were saved." + ) + return + + if message.is_stopped(): + return + if message.get_extra("agent_stop_requested"): + reply(cancelled) + return + + try: + await conversation_manager.update_conversation( + umo, + cid, + history=target_history, + token_usage=0, + ) + except Exception as update_error: + logger.error( + "Context compression storage update failed: %s.", + type(update_error).__name__, + ) + try: + stored_conversation = ( + await conversation_manager.get_conversation(umo, cid) + ) + stored_history = json.loads(stored_conversation.history) + except Exception as verify_error: + logger.error( + "Context compression storage verification failed: %s.", + type(verify_error).__name__, + ) + reply(unknown) + return + + if stored_history != target_history: + reply( + preserved if stored_history == original_history else unknown + ) + return + except Exception as error: + logger.error( + "Context compression failed before storage update: %s.", + type(error).__name__, + ) + reply(preserved) + return + + if message.get_platform_name() == "webchat": + try: + await message.send( + MessageChain( + type="agent_stats", + chain=[ + Json( + data=AgentStats( + current_context_tokens=tokens_after, + ).to_dict() + ) + ], + ) + ) + except Exception as error: + logger.warning( + "Failed to send context compression stats: %s.", + type(error).__name__, + ) + + reply("✅ Context compressed.") + async def new_conv(self, message: AstrMessageEvent) -> None: """Start a new conversation without clearing the previous history. diff --git a/astrbot/builtin_stars/builtin_commands/main.py b/astrbot/builtin_stars/builtin_commands/main.py index 6765e78cff..3143d7c0d5 100644 --- a/astrbot/builtin_stars/builtin_commands/main.py +++ b/astrbot/builtin_stars/builtin_commands/main.py @@ -63,6 +63,11 @@ async def stats(self, message: AstrMessageEvent) -> None: """Show token usage statistics for the current conversation""" await self.conversation_c.stats(message) + @filter.command("compact") + async def compact(self, message: AstrMessageEvent) -> None: + """Compress the current conversation context""" + await self.conversation_c.compact(message) + @filter.permission_type(filter.PermissionType.ADMIN) @filter.command("provider") async def provider( diff --git a/astrbot/core/agent/context/compressor.py b/astrbot/core/agent/context/compressor.py index 759604dd93..8b7c5a4a32 100644 --- a/astrbot/core/agent/context/compressor.py +++ b/astrbot/core/agent/context/compressor.py @@ -130,6 +130,7 @@ def __init__( instruction_text: str | None = None, compression_threshold: float = 0.82, token_counter: TokenCounter | None = None, + preserve_latest_round: bool = False, ) -> None: """Initialize the LLM summary compressor. @@ -139,11 +140,15 @@ 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: Token counter used to divide old and recent context. + preserve_latest_round: Whether to preserve the latest complete + user-assistant round as exact context. """ self.provider = provider 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() + self.preserve_latest_round = preserve_latest_round self.instruction_text = instruction_text or ( "Based on our full conversation history, produce a concise summary of key takeaways and/or project progress.\n" @@ -207,8 +212,15 @@ async def __call__(self, messages: list[Message]) -> list[Message]: """Use LLM to generate a summary of the conversation history. Uses round-based splitting to preserve user-assistant turn boundaries. - On LLM failure, returns the original messages unchanged (caller should - fall back to truncation). + On LLM failure, returns the original messages unchanged so the caller + can apply its configured fallback policy. + + Args: + messages: The original message list. + + Returns: + The compressed message list, or the original list when compression + cannot be completed safely. """ from .round_utils import split_into_rounds @@ -216,12 +228,40 @@ async def __call__(self, messages: list[Message]) -> list[Message]: message_rounds = [ [seg for seg in rnd if isinstance(seg, Message)] for rnd in rounds ] + latest_complete_round_index: int | None = None + if self.preserve_latest_round: + complete_round_indices = [] + for round_index, rnd in enumerate(message_rounds): + user_seen = False + for msg in rnd: + if msg.role == "user": + user_seen = True + elif user_seen and msg.role == "assistant": + complete_round_indices.append(round_index) + break + + # A summary would have no older complete round to replace. + if len(complete_round_indices) <= 1: + return messages + latest_complete_round_index = complete_round_indices[-1] + total_tokens = self.token_counter.count_tokens(messages) old_rounds, recent_rounds = self._split_recent_rounds_by_token_ratio( message_rounds, total_tokens, ) + if latest_complete_round_index is not None: + recent_start = min(len(old_rounds), latest_complete_round_index) + if not any( + msg.role != "system" + for rnd in message_rounds[:recent_start] + for msg in rnd + ): + recent_start = latest_complete_round_index + old_rounds = message_rounds[:recent_start] + recent_rounds = message_rounds[recent_start:] + # The latest user message is the active request. Keep its whole round # exact even when the ratio is 0 or the ratio budget would otherwise # summarize every round. @@ -278,7 +318,7 @@ async def __call__(self, messages: list[Message]) -> list[Message]: ) summary_content = (response.completion_text or "").strip() except Exception as e: - logger.error(f"Failed to generate summary: {e}") + logger.error("Context summary failed: %s.", type(e).__name__) return messages if not summary_content: diff --git a/astrbot/core/agent/context/config.py b/astrbot/core/agent/context/config.py index aa216d9a25..66cccb3ceb 100644 --- a/astrbot/core/agent/context/config.py +++ b/astrbot/core/agent/context/config.py @@ -33,3 +33,5 @@ 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.""" + llm_compress_preserve_latest_round: bool = False + """Whether to preserve the latest complete user-assistant round exactly.""" diff --git a/astrbot/core/agent/context/manager.py b/astrbot/core/agent/context/manager.py index 1a11ebff96..75987f3039 100644 --- a/astrbot/core/agent/context/manager.py +++ b/astrbot/core/agent/context/manager.py @@ -36,6 +36,7 @@ def __init__( keep_recent_ratio=config.llm_compress_keep_recent_ratio, instruction_text=config.llm_compress_instruction, token_counter=self.token_counter, + preserve_latest_round=config.llm_compress_preserve_latest_round, ) else: self.compressor = TruncateByTurnsCompressor( @@ -43,12 +44,18 @@ def __init__( ) async def process( - self, messages: list[Message], trusted_token_usage: int = 0 + self, + messages: list[Message], + trusted_token_usage: int = 0, + force_compress: bool = False, ) -> list[Message]: """Process the messages. Args: messages: The original message list. + trusted_token_usage: Token usage reported by the previous provider call. + force_compress: Whether to bypass automatic limits and run the configured + compressor immediately without a truncation fallback. Returns: The processed message list. @@ -57,13 +64,21 @@ async def process( result = messages # 1. 基于轮次的截断 (Enforce max turns) - if self.config.enforce_max_turns != -1: + if not force_compress and 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 force_compress: + total_tokens = self.token_counter.count_tokens(result) + return await self._run_compression( + result, + total_tokens, + allow_halving_fallback=False, + ) + # 2. 基于 token 的压缩 if self.config.max_context_tokens > 0: total_tokens = self.token_counter.count_tokens( @@ -77,18 +92,22 @@ async def process( return result except Exception as e: - logger.error(f"Error during context processing: {e}", exc_info=True) + logger.error("Context processing failed: %s.", type(e).__name__) return messages async def _run_compression( - self, messages: list[Message], prev_tokens: int + self, + messages: list[Message], + prev_tokens: int, + allow_halving_fallback: bool = True, ) -> list[Message]: - """ - Compress/truncate the messages. + """Compress or truncate the messages. Args: messages: The original message list. prev_tokens: The token count before compression. + allow_halving_fallback: Whether to halve the result if it still exceeds + the automatic compression threshold. Returns: The compressed/truncated message list. @@ -100,17 +119,25 @@ async def _run_compression( # 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}%.", - ) + if self.config.max_context_tokens > 0: + 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}%.", + ) + else: + logger.info( + f"Compress completed. {prev_tokens} -> {tokens_after_summary} tokens." + ) # last check - if self.compressor.should_compress( - messages, tokens_after_summary, self.config.max_context_tokens + if allow_halving_fallback and 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..." diff --git a/astrbot/core/astr_main_agent.py b/astrbot/core/astr_main_agent.py index 7ceb07a5a3..d0776797d3 100644 --- a/astrbot/core/astr_main_agent.py +++ b/astrbot/core/astr_main_agent.py @@ -1298,30 +1298,32 @@ def _apply_web_search_citation_prompt( req.system_prompt = f"{system_prompt}\n{WEB_SEARCH_CITATION_PROMPT}\n" -async def _get_compress_provider( - config: MainAgentBuildConfig, +async def get_context_compression_provider( + context_limit_reached_strategy: str, + llm_compress_provider_id: str, plugin_context: Context, event: AstrMessageEvent | None = None, ) -> Provider | None: """Resolve the provider used for context compression. Args: - config: Main agent build configuration. + context_limit_reached_strategy: Configured context handling strategy. + llm_compress_provider_id: Optional dedicated compression provider ID. plugin_context: Plugin context used to resolve providers. event: Optional event used for session-specific fallback selection. Returns: Compression provider, or None if compression is disabled or unavailable. """ - if config.context_limit_reached_strategy != "llm_compress": + if context_limit_reached_strategy != "llm_compress": return None - if config.llm_compress_provider_id: - provider = plugin_context.get_provider_by_id(config.llm_compress_provider_id) + if llm_compress_provider_id: + provider = plugin_context.get_provider_by_id(llm_compress_provider_id) if provider and isinstance(provider, Provider): return provider logger.warning( - "指定的上下文压缩模型 %s 不可用", - config.llm_compress_provider_id, + "Configured context compression provider %s is unavailable.", + llm_compress_provider_id, ) # fallback: use current chat provider for this session if event: @@ -1963,8 +1965,9 @@ async def build_main_agent( streaming=config.streaming_response, llm_compress_instruction=config.llm_compress_instruction, llm_compress_keep_recent_ratio=config.llm_compress_keep_recent_ratio, - llm_compress_provider=await _get_compress_provider( - config, + llm_compress_provider=await get_context_compression_provider( + config.context_limit_reached_strategy, + config.llm_compress_provider_id, plugin_context, event, ), diff --git a/tests/agent/test_context_manager.py b/tests/agent/test_context_manager.py index a596677e9b..e6ca3cf57d 100644 --- a/tests/agent/test_context_manager.py +++ b/tests/agent/test_context_manager.py @@ -70,6 +70,7 @@ def test_init_with_minimal_config(self): assert manager.token_counter is not None assert manager.truncator is not None assert manager.compressor is not None + assert not config.llm_compress_preserve_latest_round def test_init_with_llm_compressor(self): """Test initialization with LLM-based compression.""" @@ -77,6 +78,7 @@ def test_init_with_llm_compressor(self): config = ContextConfig( llm_compress_provider=mock_provider, # type: ignore llm_compress_keep_recent_ratio=0.15, + llm_compress_preserve_latest_round=True, llm_compress_instruction="Summarize the conversation", ) manager = ContextManager(config) @@ -84,6 +86,7 @@ def test_init_with_llm_compressor(self): from astrbot.core.agent.context.compressor import LLMSummaryCompressor assert isinstance(manager.compressor, LLMSummaryCompressor) + assert manager.compressor.preserve_latest_round def test_init_with_truncate_compressor(self): """Test initialization with truncate-based compression (default).""" @@ -113,6 +116,27 @@ async def test_llm_compressor_keeps_history_when_summary_is_empty(self): "LLM context compression returned an empty summary." ) + @pytest.mark.asyncio + async def test_llm_compressor_failure_log_omits_exception_details(self): + """Provider failures must not log content that could contain history.""" + from astrbot.core.agent.context.compressor import LLMSummaryCompressor + + provider = MockProvider() + provider.text_chat = AsyncMock( + side_effect=RuntimeError("secret-history from request body") + ) + compressor = LLMSummaryCompressor(provider=provider) # type: ignore[arg-type] + messages = self.create_messages(6) + + with patch("astrbot.core.agent.context.compressor.logger") as mock_logger: + result = await compressor(messages) + + assert result == messages + mock_logger.error.assert_called_once_with( + "Context summary failed: %s.", + "RuntimeError", + ) + @pytest.mark.asyncio async def test_llm_compressor_handles_textpart_content(self): from astrbot.core.agent.context.compressor import LLMSummaryCompressor @@ -299,6 +323,87 @@ async def test_llm_compressor_summarizes_system_plus_single_completed_round(self assert result[1].role == "user" assert result[2].role == "assistant" + @pytest.mark.asyncio + async def test_llm_compressor_preserves_only_complete_round_when_enabled(self): + """Do not summarize when there is no older complete round to replace.""" + from astrbot.core.agent.context.compressor import LLMSummaryCompressor + + provider = MockProvider() + compressor = LLMSummaryCompressor( + provider=provider, + keep_recent_ratio=0, + preserve_latest_round=True, + ) # type: ignore[arg-type] + messages = [ + Message(role="system", content="System prompt"), + Message(role="user", content="Question"), + Message(role="assistant", content="Answer"), + Message(role="user", content="Pending question"), + ] + + result = await compressor(messages) + + assert result == messages + assert provider.last_text_chat_kwargs is None + + @pytest.mark.asyncio + async def test_llm_compressor_preserves_latest_complete_round_when_enabled(self): + """Keep the latest completed round and later messages as exact context.""" + from astrbot.core.agent.context.compressor import LLMSummaryCompressor + + provider = MockProvider() + compressor = LLMSummaryCompressor( + provider=provider, + keep_recent_ratio=0, + preserve_latest_round=True, + ) # type: ignore[arg-type] + messages = [ + Message(role="user", content="Old question"), + Message(role="assistant", content="Old answer"), + Message(role="user", content="Latest completed question"), + Message(role="assistant", content="Latest completed answer"), + Message(role="user", content="Pending question"), + ] + + result = await compressor(messages) + + summary_contexts = provider.last_text_chat_kwargs["contexts"] + assert summary_contexts[0] == { + "role": "user", + "content": "Old question", + } + assert summary_contexts[1] == { + "role": "assistant", + "content": "Old answer", + } + assert result[-3:] == messages[-3:] + + @pytest.mark.asyncio + async def test_llm_compressor_preserves_zero_token_latest_round(self): + """A zero-token estimate cannot move the protected round into the summary.""" + from astrbot.core.agent.context.compressor import LLMSummaryCompressor + + provider = MockProvider() + compressor = LLMSummaryCompressor( + provider=provider, + keep_recent_ratio=0.15, + preserve_latest_round=True, + ) # type: ignore[arg-type] + messages = [ + Message(role="user", content="x" * 200), + Message(role="assistant", content="y" * 200), + Message(role="user", content="?"), + Message(role="assistant", content="!"), + ] + + result = await compressor(messages) + + summary_contexts = provider.last_text_chat_kwargs["contexts"] + assert summary_contexts[0] == {"role": "user", "content": "x" * 200} + assert summary_contexts[1] == {"role": "assistant", "content": "y" * 200} + assert result[-2] is messages[-2] + assert result[-1] is messages[-1] + @pytest.mark.asyncio async def test_llm_compressor_sanitizes_context_for_text_only_provider(self): from astrbot.core.agent.context.compressor import LLMSummaryCompressor @@ -546,6 +651,39 @@ async def test_trusted_usage_triggers_compression_before_provider_call(self): mock_compressor.assert_awaited_once_with(messages) assert result == compressed + @pytest.mark.asyncio + async def test_force_compression_bypasses_automatic_guards(self): + """Forced compression ignores limits, trusted usage, and truncation.""" + config = ContextConfig(max_context_tokens=0, enforce_max_turns=1) + manager = ContextManager(config) + messages = self.create_messages(6) + mock_compressor = AsyncMock(return_value=messages) + mock_compressor.should_compress = MagicMock(return_value=True) + manager.compressor = mock_compressor + manager.token_counter = MagicMock() + manager.token_counter.count_tokens.side_effect = [10, 10] + + with ( + patch.object(manager.truncator, "truncate_by_turns") as mock_turns, + patch.object(manager.truncator, "truncate_by_halving") as mock_halving, + ): + result = await manager.process( + messages, + trusted_token_usage=999, + force_compress=True, + ) + + assert result == messages + mock_compressor.assert_awaited_once_with(messages) + mock_compressor.should_compress.assert_not_called() + mock_turns.assert_not_called() + mock_halving.assert_not_called() + assert manager.token_counter.count_tokens.call_count == 2 + assert all( + call.args == (messages,) + for call in manager.token_counter.count_tokens.call_args_list + ) + @pytest.mark.asyncio async def test_token_compression_with_zero_max_tokens(self): """Test that compression is skipped when max_context_tokens is 0.""" @@ -680,8 +818,10 @@ async def test_error_handling_logs_exception(self): with patch("astrbot.core.agent.context.manager.logger") as mock_logger: result = await manager.process(messages) - # Logger error method should be called - assert mock_logger.error.called + mock_logger.error.assert_called_once_with( + "Context processing failed: %s.", + "Exception", + ) # Should return original messages on error assert result == messages diff --git a/tests/test_conversation_commands.py b/tests/test_conversation_commands.py index b0a098a4fe..8eba477a85 100644 --- a/tests/test_conversation_commands.py +++ b/tests/test_conversation_commands.py @@ -1,12 +1,142 @@ +import json from types import SimpleNamespace +from unittest.mock import AsyncMock import pytest +from astrbot.api.event import MessageEventResult from astrbot.builtin_stars.builtin_commands.commands import ( conversation as conversation_module, ) +class FakeCompactEvent: + """Minimal event implementation for manual compression command tests.""" + + def __init__( + self, + *, + group_id: str = "", + role: str = "member", + platform_name: str = "webchat", + extras: dict | None = None, + fail_stats_send: bool = False, + stopped: bool = False, + ) -> None: + self.unified_msg_origin = "webchat:private:test" + self.role = role + self.group_id = group_id + self.platform_name = platform_name + self.extras = extras or {} + self.fail_stats_send = fail_stats_send + self.stopped = stopped + self.result = None + self.sent = [] + + def get_group_id(self) -> str: + return self.group_id + + def get_platform_name(self) -> str: + return self.platform_name + + def get_extra(self, key: str, default=None): + return self.extras.get(key, default) + + def is_stopped(self) -> bool: + return self.stopped + + def plain_result(self, text: str) -> MessageEventResult: + return MessageEventResult().message(text) + + def set_result(self, result: MessageEventResult) -> None: + self.result = result + + async def send(self, chain) -> None: + self.sent.append(chain) + if self.fail_stats_send and chain.type == "agent_stats": + raise RuntimeError("stats transport failed") + + +class FakeCompressionProvider: + """Compression provider returning a configurable summary.""" + + def __init__( + self, + summary: str = "Concise summary.", + *, + fail: bool = False, + ) -> None: + self.provider_config = {} + self.summary = summary + self.fail = fail + self.call_count = 0 + + async def text_chat(self, **kwargs): + _ = kwargs + self.call_count += 1 + if self.fail: + raise RuntimeError("summary provider failed") + return SimpleNamespace(completion_text=self.summary) + + +def _compact_history() -> list[dict]: + """Build two complete rounds with a checkpoint on the latest round.""" + return [ + {"role": "user", "content": "old question " * 120}, + {"role": "assistant", "content": "old answer " * 120}, + {"role": "_checkpoint", "content": {"id": "cp-old"}}, + {"role": "user", "content": "latest question"}, + {"role": "assistant", "content": "latest answer"}, + {"role": "_checkpoint", "content": {"id": "cp-latest"}}, + ] + + +def _compact_context( + history: list[dict], + *, + settings: dict | None = None, + unique_session: bool = False, +): + """Create a command context and mocked conversation manager. + + Args: + history: Persisted conversation history. + settings: Provider setting overrides. + unique_session: Whether group conversations are isolated by member. + + Returns: + The fake command context and its conversation manager. + """ + provider_settings = { + "enable": True, + "enable_manual_context_compression": True, + "agent_runner_type": "local", + "context_limit_reached_strategy": "llm_compress", + "llm_compress_keep_recent_ratio": 0.15, + } + provider_settings.update(settings or {}) + conversation = SimpleNamespace(history=json.dumps(history)) + manager = SimpleNamespace( + get_curr_conversation_id=AsyncMock(return_value="cid-1"), + get_conversation=AsyncMock(return_value=conversation), + update_conversation=AsyncMock(), + ) + context = SimpleNamespace( + conversation_manager=manager, + get_config=lambda **kwargs: { + "provider_settings": provider_settings, + "platform_settings": {"unique_session": unique_session}, + }, + ) + return context, manager + + +def _result_text(event: FakeCompactEvent) -> str: + """Return the plain-text command result.""" + assert event.result is not None + return event.result.get_plain_text() + + @pytest.mark.asyncio async def test_clear_third_party_agent_runner_state_deletes_deerflow_thread_before_local_state( monkeypatch: pytest.MonkeyPatch, @@ -170,3 +300,411 @@ async def fake_remove_async(*args, **kwargs): "umo-3", conversation_module.DEERFLOW_THREAD_ID_KEY, ) in calls + + +@pytest.mark.asyncio +@pytest.mark.parametrize( + ("settings", "group_id", "role", "unique_session", "expected"), + [ + ({"enable": False}, "", "member", False, "AI features are disabled"), + ({"enable_manual_context_compression": False}, "", "member", False, "disabled"), + ({"agent_runner_type": "dify"}, "", "member", False, "local agent"), + ( + {"context_limit_reached_strategy": "truncate_by_turns"}, + "", + "member", + False, + "LLM context compression strategy", + ), + ({}, "group-1", "member", False, "admin permission"), + ({}, "", "member", False, "No LLM provider"), + ({}, "group-1", "member", True, "No LLM provider"), + ], +) +async def test_compact_rejects_unsupported_configuration_and_shared_group_members( + monkeypatch: pytest.MonkeyPatch, + settings: dict, + group_id: str, + role: str, + unique_session: bool, + expected: str, +): + context, manager = _compact_context( + _compact_history(), + settings=settings, + unique_session=unique_session, + ) + event = FakeCompactEvent(group_id=group_id, role=role) + provider_resolver = AsyncMock(return_value=None) + + monkeypatch.setattr( + conversation_module, + "get_context_compression_provider", + provider_resolver, + ) + + await conversation_module.ConversationCommands(context).compact(event) + + assert expected in _result_text(event) + manager.update_conversation.assert_not_awaited() + assert event.sent == [] + if expected != "No LLM provider": + provider_resolver.assert_not_awaited() + + +@pytest.mark.asyncio +@pytest.mark.parametrize( + ("platform_name", "sent_types"), + [("webchat", [None, "agent_stats"]), ("telegram", [None])], +) +async def test_compact_writes_reduced_checkpoint_history_and_expected_stats( + monkeypatch: pytest.MonkeyPatch, + platform_name: str, + sent_types: list[str | None], +): + history = _compact_history() + context, manager = _compact_context(history) + event = FakeCompactEvent(platform_name=platform_name) + provider = FakeCompressionProvider() + + async def get_provider(*args, **kwargs): + _ = args, kwargs + return provider + + monkeypatch.setattr( + conversation_module, + "get_context_compression_provider", + get_provider, + ) + + await conversation_module.ConversationCommands(context).compact(event) + + manager.update_conversation.assert_awaited_once() + update_call = manager.update_conversation.await_args + assert update_call.args == (event.unified_msg_origin, "cid-1") + saved_history = update_call.kwargs["history"] + assert update_call.kwargs["token_usage"] == 0 + assert {"role": "_checkpoint", "content": {"id": "cp-latest"}} in saved_history + assert {"role": "_checkpoint", "content": {"id": "cp-old"}} not in saved_history + assert provider.call_count == 1 + assert [chain.type for chain in event.sent] == sent_types + assert event.sent[0].get_plain_text() == "⏳ Compressing context..." + if platform_name == "webchat": + assert event.sent[1].chain[0].data["current_context_tokens"] > 0 + assert _result_text(event) == "✅ Context compressed." + + +@pytest.mark.asyncio +@pytest.mark.parametrize("provider_failure", [False, True]) +async def test_compact_does_not_write_when_summary_fails_or_is_unchanged( + monkeypatch: pytest.MonkeyPatch, + provider_failure: bool, +): + history = _compact_history() + context, manager = _compact_context(history) + event = FakeCompactEvent() + provider = FakeCompressionProvider(summary="", fail=provider_failure) + + async def get_provider(*args, **kwargs): + _ = args, kwargs + return provider + + monkeypatch.setattr( + conversation_module, + "get_context_compression_provider", + get_provider, + ) + + await conversation_module.ConversationCommands(context).compact(event) + + manager.update_conversation.assert_not_awaited() + assert "original context was preserved" in _result_text(event) + assert [chain.type for chain in event.sent] == [None] + + +@pytest.mark.asyncio +async def test_compact_reports_single_complete_round_as_not_enough_history( + monkeypatch: pytest.MonkeyPatch, +): + history = [ + {"role": "user", "content": "question"}, + {"role": "assistant", "content": "answer"}, + {"role": "_checkpoint", "content": {"id": "cp-latest"}}, + ] + context, manager = _compact_context(history) + event = FakeCompactEvent() + provider = FakeCompressionProvider() + + async def get_provider(*args, **kwargs): + _ = args, kwargs + return provider + + monkeypatch.setattr( + conversation_module, + "get_context_compression_provider", + get_provider, + ) + + await conversation_module.ConversationCommands(context).compact(event) + + manager.update_conversation.assert_not_awaited() + assert provider.call_count == 0 + assert "not enough conversation history" in _result_text(event) + + +@pytest.mark.asyncio +@pytest.mark.parametrize("conflict", ["cid", "history"]) +async def test_compact_does_not_write_when_conversation_changes( + monkeypatch: pytest.MonkeyPatch, + conflict: str, +): + history = _compact_history() + context, manager = _compact_context(history) + event = FakeCompactEvent() + provider = FakeCompressionProvider() + + if conflict == "cid": + manager.get_curr_conversation_id.side_effect = ["cid-1", "cid-1", "cid-2"] + else: + changed = SimpleNamespace( + history=json.dumps([*history, {"role": "user", "content": "new"}]) + ) + manager.get_conversation.side_effect = [ + SimpleNamespace(history=json.dumps(history)), + changed, + ] + + async def get_provider(*args, **kwargs): + _ = args, kwargs + return provider + + monkeypatch.setattr( + conversation_module, + "get_context_compression_provider", + get_provider, + ) + + await conversation_module.ConversationCommands(context).compact(event) + + manager.update_conversation.assert_not_awaited() + assert "changed during compression" in _result_text(event) + + +@pytest.mark.asyncio +async def test_compact_checks_force_stop_after_final_history_read( + monkeypatch: pytest.MonkeyPatch, +): + history = _compact_history() + context, manager = _compact_context(history) + event = FakeCompactEvent() + provider = FakeCompressionProvider() + read_count = 0 + + async def get_conversation(*args, **kwargs): + nonlocal read_count + _ = args, kwargs + read_count += 1 + if read_count == 2: + event.stopped = True + return SimpleNamespace(history=json.dumps(history)) + + async def get_provider(*args, **kwargs): + _ = args, kwargs + return provider + + manager.get_conversation.side_effect = get_conversation + monkeypatch.setattr( + conversation_module, + "get_context_compression_provider", + get_provider, + ) + + await conversation_module.ConversationCommands(context).compact(event) + + manager.update_conversation.assert_not_awaited() + assert read_count == 2 + assert event.result is None + + +@pytest.mark.asyncio +async def test_compact_does_not_write_after_stop_request( + monkeypatch: pytest.MonkeyPatch, +): + context, manager = _compact_context(_compact_history()) + event = FakeCompactEvent(extras={"agent_stop_requested": True}) + provider = FakeCompressionProvider() + + async def get_provider(*args, **kwargs): + _ = args, kwargs + return provider + + monkeypatch.setattr( + conversation_module, + "get_context_compression_provider", + get_provider, + ) + + await conversation_module.ConversationCommands(context).compact(event) + + manager.update_conversation.assert_not_awaited() + assert provider.call_count == 0 + assert "cancelled" in _result_text(event) + + +@pytest.mark.asyncio +@pytest.mark.parametrize("stop_during_summary", [False, True]) +async def test_compact_force_stop_returns_without_setting_a_result( + monkeypatch: pytest.MonkeyPatch, + stop_during_summary: bool, +): + context, manager = _compact_context(_compact_history()) + event = FakeCompactEvent(stopped=not stop_during_summary) + provider = FakeCompressionProvider() + + if stop_during_summary: + + async def stop_event_during_summary(**kwargs): + _ = kwargs + event.stopped = True + return SimpleNamespace(completion_text="Concise summary.") + + provider.text_chat = stop_event_during_summary + + async def get_provider(*args, **kwargs): + _ = args, kwargs + return provider + + monkeypatch.setattr( + conversation_module, + "get_context_compression_provider", + get_provider, + ) + + await conversation_module.ConversationCommands(context).compact(event) + + manager.update_conversation.assert_not_awaited() + assert event.result is None + assert [chain.type for chain in event.sent] == [None] + + +@pytest.mark.asyncio +@pytest.mark.parametrize( + ("verification_state", "expected"), + [ + ("target", "Context compressed"), + ("original", "original context was preserved"), + ("other", "Context state is unknown"), + ("error", "Context state is unknown"), + ], +) +async def test_compact_verifies_history_after_storage_update_error( + monkeypatch: pytest.MonkeyPatch, + caplog: pytest.LogCaptureFixture, + verification_state: str, + expected: str, +): + history = _compact_history() + context, manager = _compact_context(history) + event = FakeCompactEvent() + provider = FakeCompressionProvider() + stored_history = json.dumps(history) + read_count = 0 + + async def get_conversation(*args, **kwargs): + nonlocal read_count + _ = args, kwargs + read_count += 1 + if verification_state == "error" and read_count == 3: + raise RuntimeError("secret-history verification failure") + return SimpleNamespace(history=stored_history) + + async def update_conversation(*args, history, **kwargs): + nonlocal stored_history + _ = args, kwargs + if verification_state == "target": + stored_history = json.dumps(history) + elif verification_state == "other": + stored_history = json.dumps( + [{"role": "user", "content": "concurrent update"}] + ) + raise RuntimeError("secret-history internal readback failure") + + async def get_provider(*args, **kwargs): + _ = args, kwargs + return provider + + manager.get_conversation.side_effect = get_conversation + manager.update_conversation.side_effect = update_conversation + monkeypatch.setattr( + conversation_module, + "get_context_compression_provider", + get_provider, + ) + + await conversation_module.ConversationCommands(context).compact(event) + + manager.update_conversation.assert_awaited_once() + assert expected in _result_text(event) + if verification_state == "target": + assert [chain.type for chain in event.sent] == [None, "agent_stats"] + else: + assert [chain.type for chain in event.sent] == [None] + assert "secret-history" not in caplog.text + assert "Traceback" not in caplog.text + assert event.unified_msg_origin not in caplog.text + + +@pytest.mark.asyncio +async def test_compact_pre_update_error_log_does_not_expose_history( + monkeypatch: pytest.MonkeyPatch, + caplog: pytest.LogCaptureFixture, +): + context, manager = _compact_context(_compact_history()) + manager.get_conversation.return_value = SimpleNamespace( + history='{"secret-history": invalid}' + ) + event = FakeCompactEvent() + provider = FakeCompressionProvider() + + async def get_provider(*args, **kwargs): + _ = args, kwargs + return provider + + monkeypatch.setattr( + conversation_module, + "get_context_compression_provider", + get_provider, + ) + + await conversation_module.ConversationCommands(context).compact(event) + + manager.update_conversation.assert_not_awaited() + assert "original context was preserved" in _result_text(event) + assert "secret-history" not in caplog.text + assert "Traceback" not in caplog.text + assert event.unified_msg_origin not in caplog.text + + +@pytest.mark.asyncio +async def test_compact_stats_failure_does_not_change_success_result( + monkeypatch: pytest.MonkeyPatch, +): + context, manager = _compact_context(_compact_history()) + event = FakeCompactEvent(fail_stats_send=True) + provider = FakeCompressionProvider() + + async def get_provider(*args, **kwargs): + _ = args, kwargs + return provider + + monkeypatch.setattr( + conversation_module, + "get_context_compression_provider", + get_provider, + ) + + await conversation_module.ConversationCommands(context).compact(event) + + manager.update_conversation.assert_awaited_once() + assert [chain.type for chain in event.sent] == [None, "agent_stats"] + assert _result_text(event) == "✅ Context compressed." diff --git a/tests/unit/test_astr_main_agent.py b/tests/unit/test_astr_main_agent.py index dee8712852..4366fac340 100644 --- a/tests/unit/test_astr_main_agent.py +++ b/tests/unit/test_astr_main_agent.py @@ -582,6 +582,89 @@ async def test_select_provider_fallback_error(self, mock_event, mock_context): ) +class TestGetContextCompressionProvider: + """Tests for context compression provider resolution.""" + + @pytest.mark.asyncio + async def test_non_llm_strategy_returns_none(self, mock_event, mock_context): + """Do not resolve a provider for non-LLM compression strategies.""" + result = await ama.get_context_compression_provider( + "truncate_by_turns", + "dedicated-provider", + mock_context, + mock_event, + ) + + assert result is None + mock_context.get_provider_by_id.assert_not_called() + mock_context.get_using_provider_async.assert_not_awaited() + + @pytest.mark.asyncio + async def test_returns_valid_dedicated_provider( + self, + mock_event, + mock_context, + mock_provider, + ): + """Prefer a valid explicitly configured compression provider.""" + mock_context.get_provider_by_id.return_value = mock_provider + + result = await ama.get_context_compression_provider( + "llm_compress", + "dedicated-provider", + mock_context, + mock_event, + ) + + assert result is mock_provider + mock_context.get_provider_by_id.assert_called_once_with("dedicated-provider") + mock_context.get_using_provider_async.assert_not_awaited() + + @pytest.mark.asyncio + async def test_invalid_dedicated_provider_falls_back_by_event_umo( + self, + mock_event, + mock_context, + mock_provider, + ): + """Use the event-scoped chat provider when the dedicated one is invalid.""" + mock_context.get_provider_by_id.return_value = "not-a-provider" + mock_context.get_using_provider_async = AsyncMock(return_value=mock_provider) + + result = await ama.get_context_compression_provider( + "llm_compress", + "invalid-provider", + mock_context, + mock_event, + ) + + assert result is mock_provider + mock_context.get_provider_by_id.assert_called_once_with("invalid-provider") + mock_context.get_using_provider_async.assert_awaited_once_with( + umo=mock_event.unified_msg_origin + ) + + @pytest.mark.asyncio + async def test_fallback_value_error_returns_none(self, mock_event, mock_context): + """Treat an invalid event-scoped fallback provider as unavailable.""" + mock_context.get_using_provider_async = AsyncMock( + side_effect=ValueError("invalid provider type") + ) + + result = await ama.get_context_compression_provider( + "llm_compress", + "", + mock_context, + mock_event, + ) + + assert result is None + mock_context.get_provider_by_id.assert_not_called() + mock_context.get_using_provider_async.assert_awaited_once_with( + umo=mock_event.unified_msg_origin + ) + + @pytest.mark.asyncio async def test_provider_manager_async_selection_uses_session_preference(monkeypatch): preferred_provider = object() From fff153d5ab48d2f241be2b5618db4579992035ea Mon Sep 17 00:00:00 2001 From: C10H14N2O5 <100066858+C10H14N2O5@users.noreply.github.com> Date: Mon, 24 Aug 2026 17:50:54 +0800 Subject: [PATCH 02/16] fix: keep compact progress transient in webchat --- .../builtin_commands/commands/conversation.py | 7 +++- astrbot/dashboard/services/chat_service.py | 16 ++++++-- .../dashboard/services/live_chat_service.py | 7 +++- .../dashboard/services/open_api_service.py | 9 ++++- tests/test_chat_route.py | 38 +++++++++++++++---- tests/test_conversation_commands.py | 18 ++++++--- tests/unit/test_live_chat_service.py | 28 ++++++++++++++ 7 files changed, 103 insertions(+), 20 deletions(-) diff --git a/astrbot/builtin_stars/builtin_commands/commands/conversation.py b/astrbot/builtin_stars/builtin_commands/commands/conversation.py index 3443f64bc2..d6d20e4535 100644 --- a/astrbot/builtin_stars/builtin_commands/commands/conversation.py +++ b/astrbot/builtin_stars/builtin_commands/commands/conversation.py @@ -245,7 +245,12 @@ def reply(text: str) -> None: reply("❌ No LLM provider is available for context compression.") return - await message.send(MessageChain().message("⏳ Compressing context...")) + progress_type = ( + "webchat_ephemeral" if message.get_platform_name() == "webchat" else None + ) + await message.send( + MessageChain(type=progress_type).message("⏳ Compressing context...") + ) try: async with session_lock_manager.acquire_lock(umo): diff --git a/astrbot/dashboard/services/chat_service.py b/astrbot/dashboard/services/chat_service.py index 54ec00cb2b..da8241fdcf 100644 --- a/astrbot/dashboard/services/chat_service.py +++ b/astrbot/dashboard/services/chat_service.py @@ -41,6 +41,7 @@ SSE_HEARTBEAT = ": heartbeat\n\n" CHAT_RUN_SUBSCRIBER_QUEUE_SIZE = 256 +WEBCHAT_EPHEMERAL_CHAIN_TYPE = "webchat_ephemeral" # Uploaded chat attachments larger than this are rejected. MAX_UPLOAD_FILE_SIZE_MB = 512 MAX_UPLOAD_FILE_SIZE_BYTES = MAX_UPLOAD_FILE_SIZE_MB * 1024 * 1024 @@ -1144,8 +1145,13 @@ async def flush_pending_bot_message(): attachment_saved_payload = None if msg_type == "plain": - for accumulator in (pending_accumulator, display_accumulator): - accumulator.add_plain( + display_accumulator.add_plain( + result_text, + chain_type=chain_type, + streaming=streaming, + ) + if chain_type != WEBCHAT_EPHEMERAL_CHAIN_TYPE: + pending_accumulator.add_plain( result_text, chain_type=chain_type, streaming=streaming, @@ -1193,7 +1199,11 @@ async def flush_pending_bot_message(): or pending_agent_stats ) elif (streaming and msg_type == "complete") or not streaming: - if chain_type not in ("tool_call", "tool_call_result"): + if chain_type not in ( + "tool_call", + "tool_call_result", + WEBCHAT_EPHEMERAL_CHAIN_TYPE, + ): should_save = True if should_save: diff --git a/astrbot/dashboard/services/live_chat_service.py b/astrbot/dashboard/services/live_chat_service.py index 16b7eed0ad..0a1cfc9abb 100644 --- a/astrbot/dashboard/services/live_chat_service.py +++ b/astrbot/dashboard/services/live_chat_service.py @@ -30,6 +30,7 @@ from astrbot.core.utils.astrbot_path import get_astrbot_data_path, get_astrbot_temp_path from astrbot.core.utils.datetime_utils import generate_timestamp_id, to_utc_isoformat from astrbot.dashboard.services.chat_service import ( + WEBCHAT_EPHEMERAL_CHAIN_TYPE, BotMessageAccumulator, build_bot_history_content, collect_plain_text_from_message_parts, @@ -704,7 +705,10 @@ async def send_attachment_saved_event(part: dict | None) -> None: outgoing = {"ct": "chat", **result} await self.send_chat_payload(session, outgoing, send_json) - if result_type == "plain": + if ( + result_type == "plain" + and chain_type != WEBCHAT_EPHEMERAL_CHAIN_TYPE + ): message_accumulator.add_plain( result_text, chain_type=chain_type, @@ -755,6 +759,7 @@ async def send_attachment_saved_event(part: dict | None) -> None: "tool_call", "tool_call_result", "agent_stats", + WEBCHAT_EPHEMERAL_CHAIN_TYPE, ): should_save = True diff --git a/astrbot/dashboard/services/open_api_service.py b/astrbot/dashboard/services/open_api_service.py index 131b8e3b98..822f19f1c9 100644 --- a/astrbot/dashboard/services/open_api_service.py +++ b/astrbot/dashboard/services/open_api_service.py @@ -27,6 +27,7 @@ DEFAULT_OPEN_API_SCOPES, ) from astrbot.dashboard.services.chat_service import ( + WEBCHAT_EPHEMERAL_CHAIN_TYPE, BotMessageAccumulator, collect_plain_text_from_message_parts, ) @@ -487,7 +488,7 @@ async def handle_chat_ws_send( await send_json(result) - if msg_type == "plain": + if msg_type == "plain" and chain_type != WEBCHAT_EPHEMERAL_CHAIN_TYPE: message_accumulator.add_plain( result_text, chain_type=chain_type, @@ -507,7 +508,11 @@ async def handle_chat_ws_send( message_accumulator.has_content() or refs or agent_stats ) elif (streaming and msg_type == "complete") or not streaming: - if chain_type not in ("tool_call", "tool_call_result"): + if chain_type not in ( + "tool_call", + "tool_call_result", + WEBCHAT_EPHEMERAL_CHAIN_TYPE, + ): should_save = True if should_save: diff --git a/tests/test_chat_route.py b/tests/test_chat_route.py index 507f3157ba..41d0620f26 100644 --- a/tests/test_chat_route.py +++ b/tests/test_chat_route.py @@ -152,17 +152,36 @@ async def test_chat_stream_disconnect_does_not_own_run_lifecycle( run.run_id, { "type": "plain", - "data": "completed after refresh", - "streaming": True, + "data": "⏳ Compressing context...", + "streaming": False, + "chain_type": "webchat_ephemeral", "message_id": run.run_id, }, ) + for _ in range(10): + if run.message_parts: + break + await asyncio.sleep(0) + assert run.message_parts == [ + {"type": "plain", "text": "⏳ Compressing context..."} + ] + await chat_service.webchat_queue_mgr.put_back_queue( run.run_id, { - "type": "complete", - "data": "completed after refresh", - "streaming": True, + "type": "plain", + "data": json.dumps({"current_context_tokens": 42}), + "streaming": False, + "chain_type": "agent_stats", + "message_id": run.run_id, + }, + ) + await chat_service.webchat_queue_mgr.put_back_queue( + run.run_id, + { + "type": "plain", + "data": "✅ Context compressed.", + "streaming": False, "message_id": run.run_id, }, ) @@ -177,8 +196,13 @@ async def test_chat_stream_disconnect_does_not_own_run_lifecycle( ) await asyncio.wait_for(run.task, timeout=1) - saved_parts = service.save_bot_message.await_args.args[1] - assert saved_parts == [{"type": "plain", "text": "completed after refresh"}] + service.save_bot_message.assert_awaited_once() + save_args = service.save_bot_message.await_args.args + assert save_args[1] == [{"type": "plain", "text": "✅ Context compressed."}] + assert save_args[2] == {"current_context_tokens": 42} + assert run.message_parts == [ + {"type": "plain", "text": "✅ Context compressed."} + ] assert run.run_id not in service.chat_runs finally: if run.task and not run.task.done(): diff --git a/tests/test_conversation_commands.py b/tests/test_conversation_commands.py index 8eba477a85..bb8bbd59ae 100644 --- a/tests/test_conversation_commands.py +++ b/tests/test_conversation_commands.py @@ -355,7 +355,7 @@ async def test_compact_rejects_unsupported_configuration_and_shared_group_member @pytest.mark.asyncio @pytest.mark.parametrize( ("platform_name", "sent_types"), - [("webchat", [None, "agent_stats"]), ("telegram", [None])], + [("webchat", ["webchat_ephemeral", "agent_stats"]), ("telegram", [None])], ) async def test_compact_writes_reduced_checkpoint_history_and_expected_stats( monkeypatch: pytest.MonkeyPatch, @@ -419,7 +419,7 @@ async def get_provider(*args, **kwargs): manager.update_conversation.assert_not_awaited() assert "original context was preserved" in _result_text(event) - assert [chain.type for chain in event.sent] == [None] + assert [chain.type for chain in event.sent] == ["webchat_ephemeral"] @pytest.mark.asyncio @@ -584,7 +584,7 @@ async def get_provider(*args, **kwargs): manager.update_conversation.assert_not_awaited() assert event.result is None - assert [chain.type for chain in event.sent] == [None] + assert [chain.type for chain in event.sent] == ["webchat_ephemeral"] @pytest.mark.asyncio @@ -646,9 +646,12 @@ async def get_provider(*args, **kwargs): manager.update_conversation.assert_awaited_once() assert expected in _result_text(event) if verification_state == "target": - assert [chain.type for chain in event.sent] == [None, "agent_stats"] + assert [chain.type for chain in event.sent] == [ + "webchat_ephemeral", + "agent_stats", + ] else: - assert [chain.type for chain in event.sent] == [None] + assert [chain.type for chain in event.sent] == ["webchat_ephemeral"] assert "secret-history" not in caplog.text assert "Traceback" not in caplog.text assert event.unified_msg_origin not in caplog.text @@ -706,5 +709,8 @@ async def get_provider(*args, **kwargs): await conversation_module.ConversationCommands(context).compact(event) manager.update_conversation.assert_awaited_once() - assert [chain.type for chain in event.sent] == [None, "agent_stats"] + assert [chain.type for chain in event.sent] == [ + "webchat_ephemeral", + "agent_stats", + ] assert _result_text(event) == "✅ Context compressed." diff --git a/tests/unit/test_live_chat_service.py b/tests/unit/test_live_chat_service.py index f7c25acc9f..5cb3035689 100644 --- a/tests/unit/test_live_chat_service.py +++ b/tests/unit/test_live_chat_service.py @@ -253,6 +253,9 @@ async def test_handle_chat_message_scopes_events_to_request_by_default(): return_value=[{"type": "plain", "text": "hello"}] ) service.ensure_chat_subscription = AsyncMock(return_value="subscription-1") + service.save_bot_message = AsyncMock( + return_value=SimpleNamespace(id=2, created_at=datetime.now(UTC)) + ) async def send_json(payload: dict) -> None: sent.append(payload) @@ -274,6 +277,25 @@ async def send_json(payload: dict) -> None: input_queue = webchat_queue_mgr.get_or_create_queue(session_id) await asyncio.wait_for(input_queue.get(), timeout=1) + await webchat_queue_mgr.put_back_queue( + message_id, + { + "type": "plain", + "data": "⏳ Compressing context...", + "streaming": False, + "chain_type": "webchat_ephemeral", + "message_id": message_id, + }, + ) + await webchat_queue_mgr.put_back_queue( + message_id, + { + "type": "plain", + "data": "✅ Context compressed.", + "streaming": False, + "message_id": message_id, + }, + ) await webchat_queue_mgr.put_back_queue( message_id, { @@ -289,6 +311,12 @@ async def send_json(payload: dict) -> None: assert sent[0]["message_id"] == message_id assert sent[-1]["type"] == "end" assert sent[-1]["message_id"] == message_id + assert [ + payload["data"] for payload in sent if payload.get("type") == "plain" + ] == ["⏳ Compressing context...", "✅ Context compressed."] + service.save_bot_message.assert_awaited_once() + saved_parts = service.save_bot_message.await_args.args[1] + assert saved_parts == [{"type": "plain", "text": "✅ Context compressed."}] finally: if not task.done(): task.cancel() From e2f40b3425ced1ba8f5b1aa1ef980fbaf0029b06 Mon Sep 17 00:00:00 2001 From: C10H14N2O5 <100066858+C10H14N2O5@users.noreply.github.com> Date: Mon, 24 Aug 2026 18:09:45 +0800 Subject: [PATCH 03/16] fix: avoid logging context token counts --- astrbot/core/agent/context/manager.py | 14 +------------- 1 file changed, 1 insertion(+), 13 deletions(-) diff --git a/astrbot/core/agent/context/manager.py b/astrbot/core/agent/context/manager.py index 75987f3039..be77ca5b16 100644 --- a/astrbot/core/agent/context/manager.py +++ b/astrbot/core/agent/context/manager.py @@ -119,19 +119,7 @@ async def _run_compression( # double check tokens_after_summary = self.token_counter.count_tokens(messages) - if self.config.max_context_tokens > 0: - 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}%.", - ) - else: - logger.info( - f"Compress completed. {prev_tokens} -> {tokens_after_summary} tokens." - ) + logger.info("Compress completed.") # last check if allow_halving_fallback and self.compressor.should_compress( From 1f239b4ee4723b1df038aec8a4823b14ba4cfca2 Mon Sep 17 00:00:00 2001 From: C10H14N2O5 <100066858+C10H14N2O5@users.noreply.github.com> Date: Mon, 24 Aug 2026 18:18:30 +0800 Subject: [PATCH 04/16] fix: preserve context compression token metrics --- astrbot/core/agent/context/manager.py | 20 +++++++++++++++---- astrbot/core/agent/context/token_counter.py | 10 +++++----- .../agent/runners/tool_loop_agent_runner.py | 2 +- tests/agent/test_context_manager.py | 8 ++++---- tests/agent/test_token_counter.py | 6 +++--- tests/test_tool_loop_agent_runner.py | 4 ++-- 6 files changed, 31 insertions(+), 19 deletions(-) diff --git a/astrbot/core/agent/context/manager.py b/astrbot/core/agent/context/manager.py index be77ca5b16..d8e29fc900 100644 --- a/astrbot/core/agent/context/manager.py +++ b/astrbot/core/agent/context/manager.py @@ -46,14 +46,14 @@ def __init__( async def process( self, messages: list[Message], - trusted_token_usage: int = 0, + reported_token_usage: int = 0, force_compress: bool = False, ) -> list[Message]: """Process the messages. Args: messages: The original message list. - trusted_token_usage: Token usage reported by the previous provider call. + reported_token_usage: Token usage reported by the previous provider call. force_compress: Whether to bypass automatic limits and run the configured compressor immediately without a truncation fallback. @@ -82,7 +82,7 @@ async def process( # 2. 基于 token 的压缩 if self.config.max_context_tokens > 0: total_tokens = self.token_counter.count_tokens( - result, trusted_token_usage + result, reported_token_usage ) if self.compressor.should_compress( @@ -119,7 +119,19 @@ async def _run_compression( # double check tokens_after_summary = self.token_counter.count_tokens(messages) - logger.info("Compress completed.") + if self.config.max_context_tokens > 0: + 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}%.", + ) + else: + logger.info( + f"Compress completed. {prev_tokens} -> {tokens_after_summary} tokens." + ) # last check if allow_halving_fallback and self.compressor.should_compress( diff --git a/astrbot/core/agent/context/token_counter.py b/astrbot/core/agent/context/token_counter.py index 31c274af64..696e219959 100644 --- a/astrbot/core/agent/context/token_counter.py +++ b/astrbot/core/agent/context/token_counter.py @@ -13,13 +13,13 @@ class TokenCounter(Protocol): """ def count_tokens( - self, messages: list[Message], trusted_token_usage: int = 0 + self, messages: list[Message], reported_token_usage: int = 0 ) -> int: """Count the total tokens in the message list. Args: messages: The message list. - trusted_token_usage: The total token usage that LLM API returned. + reported_token_usage: The total token usage that LLM API returned. For some cases, this value is more accurate. But some API does not return it, so the value defaults to 0. @@ -54,10 +54,10 @@ class EstimateTokenCounter: """ def count_tokens( - self, messages: list[Message], trusted_token_usage: int = 0 + self, messages: list[Message], reported_token_usage: int = 0 ) -> int: - if trusted_token_usage > 0: - return trusted_token_usage + if reported_token_usage > 0: + return reported_token_usage total = 0 for msg in messages: diff --git a/astrbot/core/agent/runners/tool_loop_agent_runner.py b/astrbot/core/agent/runners/tool_loop_agent_runner.py index ed5bb3df62..e473d4af54 100644 --- a/astrbot/core/agent/runners/tool_loop_agent_runner.py +++ b/astrbot/core/agent/runners/tool_loop_agent_runner.py @@ -829,7 +829,7 @@ async def step(self): processed_messages = await self._await_or_stop( self.request_context_manager.process( self.run_context.messages, - trusted_token_usage=token_usage, + reported_token_usage=token_usage, ) ) if processed_messages is None: diff --git a/tests/agent/test_context_manager.py b/tests/agent/test_context_manager.py index e6ca3cf57d..2ccfb51555 100644 --- a/tests/agent/test_context_manager.py +++ b/tests/agent/test_context_manager.py @@ -635,7 +635,7 @@ def mock_should_compress(*args, **kwargs): assert len(result) <= len(messages) @pytest.mark.asyncio - async def test_trusted_usage_triggers_compression_before_provider_call(self): + async def test_reported_usage_triggers_compression_before_provider_call(self): config = ContextConfig(max_context_tokens=100, truncate_turns=1) manager = ContextManager(config) messages = [self.create_message("user", "short")] @@ -644,7 +644,7 @@ async def test_trusted_usage_triggers_compression_before_provider_call(self): mock_compressor.should_compress = MagicMock(side_effect=[True, False]) manager.compressor = mock_compressor - result = await manager.process(messages, trusted_token_usage=83) + result = await manager.process(messages, reported_token_usage=83) first_check = mock_compressor.should_compress.call_args_list[0] assert first_check.args == (messages, 83, 100) @@ -653,7 +653,7 @@ async def test_trusted_usage_triggers_compression_before_provider_call(self): @pytest.mark.asyncio async def test_force_compression_bypasses_automatic_guards(self): - """Forced compression ignores limits, trusted usage, and truncation.""" + """Forced compression ignores limits, reported usage, and truncation.""" config = ContextConfig(max_context_tokens=0, enforce_max_turns=1) manager = ContextManager(config) messages = self.create_messages(6) @@ -669,7 +669,7 @@ async def test_force_compression_bypasses_automatic_guards(self): ): result = await manager.process( messages, - trusted_token_usage=999, + reported_token_usage=999, force_compress=True, ) diff --git a/tests/agent/test_token_counter.py b/tests/agent/test_token_counter.py index 0b72403e1b..6375b45452 100644 --- a/tests/agent/test_token_counter.py +++ b/tests/agent/test_token_counter.py @@ -127,8 +127,8 @@ def test_plain_text_estimate_unchanged(self): assert counter.count_tokens([_msg("user", "你" * 100)]) == 60 -class TestTrustedUsage: - def test_trusted_overrides(self): +class TestReportedUsage: + def test_reported_overrides(self): """如果 API 返回了 token 数,直接用它不做估算。""" msg = _msg( "user", @@ -139,7 +139,7 @@ def test_trusted_overrides(self): ), ], ) - tokens = counter.count_tokens([msg], trusted_token_usage=42) + tokens = counter.count_tokens([msg], reported_token_usage=42) assert tokens == 42 diff --git a/tests/test_tool_loop_agent_runner.py b/tests/test_tool_loop_agent_runner.py index ef8e818bb5..03e65bc928 100644 --- a/tests/test_tool_loop_agent_runner.py +++ b/tests/test_tool_loop_agent_runner.py @@ -573,7 +573,7 @@ async def test_max_step_final_request_includes_limit_prompt( streaming=False, ) - async def snapshot_context_manager(messages, trusted_token_usage=0): + async def snapshot_context_manager(messages, reported_token_usage=0): return list(messages) runner.request_context_manager.process = snapshot_context_manager @@ -603,7 +603,7 @@ async def test_tool_loop_next_request_includes_tool_result( streaming=False, ) - async def snapshot_context_manager(messages, trusted_token_usage=0): + async def snapshot_context_manager(messages, reported_token_usage=0): return list(messages) runner.request_context_manager.process = snapshot_context_manager From e887c44367ec06e76f86a9a4d94c08ac8db2747c Mon Sep 17 00:00:00 2001 From: C10H14N2O5 <100066858+C10H14N2O5@users.noreply.github.com> Date: Mon, 24 Aug 2026 19:20:21 +0800 Subject: [PATCH 05/16] test: preserve streaming disconnect regression coverage --- tests/test_chat_route.py | 5 +++++ 1 file changed, 5 insertions(+) diff --git a/tests/test_chat_route.py b/tests/test_chat_route.py index 41d0620f26..5166c15e0e 100644 --- a/tests/test_chat_route.py +++ b/tests/test_chat_route.py @@ -277,6 +277,11 @@ async def test_resumed_stream_starts_with_full_snapshot(chat_service_instance): ): await chat_service.webchat_queue_mgr.put_back_queue(run.run_id, payload) await asyncio.wait_for(run.task, timeout=1) + service.save_bot_message.assert_awaited_once() + saved_parts = service.save_bot_message.await_args.args[1] + assert saved_parts == [ + {"type": "plain", "text": "before refresh and after refresh"}, + ] finally: if run.task and not run.task.done(): run.task.cancel() From 4b8bfc9a724afdae73edaaffd59d981a854667f3 Mon Sep 17 00:00:00 2001 From: C10H14N2O5 <100066858+C10H14N2O5@users.noreply.github.com> Date: Sun, 30 Aug 2026 19:06:19 +0800 Subject: [PATCH 06/16] refactor: adapt manual compression to embedded runner config --- .../builtin_commands/commands/conversation.py | 28 ++- astrbot/core/config/agent_runner.py | 1 + astrbot/core/config/default.py | 9 + astrbot/core/utils/migra_helper.py | 4 + .../en-US/features/config-metadata.json | 4 + .../ru-RU/features/config-metadata.json | 4 + .../zh-CN/features/config-metadata.json | 4 + tests/test_conversation_commands.py | 228 +++++++++++++++--- tests/unit/test_agent_runner_config.py | 27 ++- tests/unit/test_config.py | 46 +++- tests/unit/test_config_profile_service.py | 15 +- tests/unit/test_live_chat_service.py | 22 +- tests/unit/test_open_api_service_ws.py | 99 ++++++++ 13 files changed, 432 insertions(+), 59 deletions(-) diff --git a/astrbot/builtin_stars/builtin_commands/commands/conversation.py b/astrbot/builtin_stars/builtin_commands/commands/conversation.py index d6d20e4535..393a7d09bb 100644 --- a/astrbot/builtin_stars/builtin_commands/commands/conversation.py +++ b/astrbot/builtin_stars/builtin_commands/commands/conversation.py @@ -196,6 +196,7 @@ def reply(text: str) -> None: umo = message.unified_msg_origin cfg = self.context.get_config(umo=umo) provider_settings = cfg.get("provider_settings", {}) + agent_runner = cfg.get("agent_runner", {}) conversation_manager = self.context.conversation_manager is_unique_session = cfg.get("platform_settings", {}).get( @@ -214,18 +215,20 @@ def reply(text: str) -> None: reply("❌ AI features are disabled for this session.") return - if not provider_settings.get("enable_manual_context_compression", False): + if agent_runner.get("runner_type") != "local": + reply("❌ /compact is supported only by the local agent runner.") + return + + runner_config = agent_runner.get("config", {}) + compression_config = runner_config.get("compression", {}) + if not compression_config.get("enable_manual_context_compression", False): reply( "❌ Manual context compression is disabled. Enable it in Context " "Management first." ) return - if provider_settings.get("agent_runner_type", "local") != "local": - reply("❌ /compact is supported only by the local agent runner.") - return - - strategy = provider_settings.get("context_limit_reached_strategy") + strategy = compression_config.get("overflow_strategy") if strategy != "llm_compress": reply("❌ /compact requires the LLM context compression strategy.") return @@ -237,7 +240,7 @@ def reply(text: str) -> None: compression_provider = await get_context_compression_provider( strategy, - provider_settings.get("llm_compress_provider_id", ""), + compression_config.get("provider_id", ""), self.context, message, ) @@ -291,11 +294,9 @@ def reply(text: str) -> None: context_manager = ContextManager( ContextConfig( - llm_compress_instruction=provider_settings.get( - "llm_compress_instruction" - ), - llm_compress_keep_recent_ratio=provider_settings.get( - "llm_compress_keep_recent_ratio", + llm_compress_instruction=compression_config.get("instruction"), + llm_compress_keep_recent_ratio=compression_config.get( + "keep_recent_ratio", 0.15, ), llm_compress_preserve_latest_round=True, @@ -377,6 +378,9 @@ def reply(text: str) -> None: reply(preserved) return + if message.is_stopped(): + return + if message.get_platform_name() == "webchat": try: await message.send( diff --git a/astrbot/core/config/agent_runner.py b/astrbot/core/config/agent_runner.py index b8c04894cc..fbb766e333 100644 --- a/astrbot/core/config/agent_runner.py +++ b/astrbot/core/config/agent_runner.py @@ -22,6 +22,7 @@ "max_turns": -1, "trim_turns": 1, "overflow_strategy": "llm_compress", + "enable_manual_context_compression": False, "instruction": "", "keep_recent_ratio": 0.15, "provider_id": "", diff --git a/astrbot/core/config/default.py b/astrbot/core/config/default.py index 549aa55799..235ef4ec73 100644 --- a/astrbot/core/config/default.py +++ b/astrbot/core/config/default.py @@ -4056,6 +4056,15 @@ def get_local_permission_defaults(system: str | None = None) -> dict: }, "hint": "普通会话历史仅在超过“压缩前最多保留对话轮数”后执行该策略;请求发送前也会在上下文 token 接近模型窗口时使用同一策略保护本次请求。", }, + "agent_runner.config.compression.enable_manual_context_compression": { + "description": "手动上下文压缩(实验性)", + "type": "bool", + "hint": "启用后,/compact 将使用 LLM 摘要当前上下文。摘要可能遗漏细节、角色状态或叙事事实;压缩失败时将保留原历史。", + "condition": { + "agent_runner.config.compression.overflow_strategy": "llm_compress", + "agent_runner.runner_type": "local", + }, + }, "agent_runner.config.compression.instruction": { "description": "上下文压缩提示词", "type": "text", diff --git a/astrbot/core/utils/migra_helper.py b/astrbot/core/utils/migra_helper.py index 86e17d962e..060dac10af 100644 --- a/astrbot/core/utils/migra_helper.py +++ b/astrbot/core/utils/migra_helper.py @@ -40,6 +40,7 @@ "tool_call_timeout", "sanitize_context_by_modalities", "context_limit_reached_strategy", + "enable_manual_context_compression", "llm_compress_instruction", "llm_compress_keep_recent_ratio", "llm_compress_provider_id", @@ -222,6 +223,9 @@ def _migrate_agent_runner_config( "overflow_strategy": provider_settings.get( "context_limit_reached_strategy", "llm_compress" ), + "enable_manual_context_compression": provider_settings.get( + "enable_manual_context_compression", False + ), "instruction": provider_settings.get("llm_compress_instruction", ""), "keep_recent_ratio": provider_settings.get( "llm_compress_keep_recent_ratio", 0.15 diff --git a/dashboard/src/i18n/locales/en-US/features/config-metadata.json b/dashboard/src/i18n/locales/en-US/features/config-metadata.json index aa12b9efa2..90bc5d5437 100644 --- a/dashboard/src/i18n/locales/en-US/features/config-metadata.json +++ b/dashboard/src/i18n/locales/en-US/features/config-metadata.json @@ -378,6 +378,10 @@ ], "hint": "Persistent conversation history uses this strategy only after exceeding 'Max Turns Before Compression'. Before each request, the same strategy may also protect the in-flight context when tokens approach the model window." }, + "enable_manual_context_compression": { + "description": "Manual Context Compression (Experimental)", + "hint": "When enabled, /compact uses an LLM to summarize the current context. The summary may omit details, role state, or narrative facts. If compression fails, the original history is preserved." + }, "instruction": { "description": "Context Compression Instruction", "hint": "If empty, the default prompt will be used." diff --git a/dashboard/src/i18n/locales/ru-RU/features/config-metadata.json b/dashboard/src/i18n/locales/ru-RU/features/config-metadata.json index c062df0eec..b63e517e1d 100644 --- a/dashboard/src/i18n/locales/ru-RU/features/config-metadata.json +++ b/dashboard/src/i18n/locales/ru-RU/features/config-metadata.json @@ -507,6 +507,10 @@ ], "hint": "Постоянная история диалога использует эту стратегию только после превышения лимита раундов. Перед каждым запросом та же стратегия может защищать текущий контекст, когда токены приближаются к окну модели." }, + "enable_manual_context_compression": { + "description": "Ручное сжатие контекста (экспериментальная функция)", + "hint": "После включения команда /compact использует LLM для создания краткого содержания текущего контекста. Сводка может упустить детали, состояние роли или факты повествования. При сбое исходная история сохраняется." + }, "instruction": { "description": "Инструкция для сжатия контекста", "hint": "Если пусто, используется промпт по умолчанию." diff --git a/dashboard/src/i18n/locales/zh-CN/features/config-metadata.json b/dashboard/src/i18n/locales/zh-CN/features/config-metadata.json index fc105671d3..e618f50f62 100644 --- a/dashboard/src/i18n/locales/zh-CN/features/config-metadata.json +++ b/dashboard/src/i18n/locales/zh-CN/features/config-metadata.json @@ -393,6 +393,10 @@ "labels": ["按对话轮数截断", "由 LLM 压缩上下文"], "hint": "普通会话历史仅在超过\"压缩前最多保留对话轮数\"后执行该策略;请求发送前也会在上下文 token 接近模型窗口时使用同一策略保护本次请求。" }, + "enable_manual_context_compression": { + "description": "手动上下文压缩(实验性)", + "hint": "启用后,/compact 将使用 LLM 摘要当前上下文。摘要可能遗漏细节、角色状态或叙事事实;压缩失败时将保留原历史。" + }, "instruction": { "description": "上下文压缩提示词", "hint": "如果为空则使用默认提示词。" diff --git a/tests/test_conversation_commands.py b/tests/test_conversation_commands.py index bb8bbd59ae..eca21627c2 100644 --- a/tests/test_conversation_commands.py +++ b/tests/test_conversation_commands.py @@ -94,27 +94,33 @@ def _compact_history() -> list[dict]: def _compact_context( history: list[dict], *, - settings: dict | None = None, + provider_settings: dict | None = None, + runner_type: str = "local", + compression_settings: dict | None = None, unique_session: bool = False, ): """Create a command context and mocked conversation manager. Args: history: Persisted conversation history. - settings: Provider setting overrides. + provider_settings: Provider setting overrides. + runner_type: Embedded Agent Runner type. + compression_settings: Local compression setting overrides. unique_session: Whether group conversations are isolated by member. Returns: The fake command context and its conversation manager. """ - provider_settings = { - "enable": True, + effective_provider_settings = {"enable": True} + effective_provider_settings.update(provider_settings or {}) + compression_config = { "enable_manual_context_compression": True, - "agent_runner_type": "local", - "context_limit_reached_strategy": "llm_compress", - "llm_compress_keep_recent_ratio": 0.15, + "overflow_strategy": "llm_compress", + "instruction": "Keep the important context.", + "keep_recent_ratio": 0.15, + "provider_id": "compression-provider", } - provider_settings.update(settings or {}) + compression_config.update(compression_settings or {}) conversation = SimpleNamespace(history=json.dumps(history)) manager = SimpleNamespace( get_curr_conversation_id=AsyncMock(return_value="cid-1"), @@ -124,7 +130,15 @@ def _compact_context( context = SimpleNamespace( conversation_manager=manager, get_config=lambda **kwargs: { - "provider_settings": provider_settings, + "provider_settings": effective_provider_settings, + "agent_runner": { + "runner_type": runner_type, + "config": ( + {"compression": compression_config} + if runner_type == "local" + else {} + ), + }, "platform_settings": {"unique_session": unique_session}, }, ) @@ -304,26 +318,54 @@ async def fake_remove_async(*args, **kwargs): @pytest.mark.asyncio @pytest.mark.parametrize( - ("settings", "group_id", "role", "unique_session", "expected"), + ( + "provider_settings", + "runner_type", + "compression_settings", + "group_id", + "role", + "unique_session", + "expected", + ), [ - ({"enable": False}, "", "member", False, "AI features are disabled"), - ({"enable_manual_context_compression": False}, "", "member", False, "disabled"), - ({"agent_runner_type": "dify"}, "", "member", False, "local agent"), ( - {"context_limit_reached_strategy": "truncate_by_turns"}, + {"enable": False}, + "local", + {}, + "", + "member", + False, + "AI features are disabled", + ), + ( + {}, + "local", + {"enable_manual_context_compression": False}, + "", + "member", + False, + "disabled", + ), + ({}, "dify", {}, "", "member", False, "local agent"), + ( + {}, + "local", + {"overflow_strategy": "truncate_by_turns"}, "", "member", False, "LLM context compression strategy", ), - ({}, "group-1", "member", False, "admin permission"), - ({}, "", "member", False, "No LLM provider"), - ({}, "group-1", "member", True, "No LLM provider"), + ({}, "local", {}, "group-1", "member", False, "admin permission"), + ({}, "local", {}, "", "member", False, "No LLM provider"), + ({}, "local", {}, "group-1", "member", True, "No LLM provider"), ], ) async def test_compact_rejects_unsupported_configuration_and_shared_group_members( monkeypatch: pytest.MonkeyPatch, - settings: dict, + provider_settings: dict, + runner_type: str, + compression_settings: dict, group_id: str, role: str, unique_session: bool, @@ -331,7 +373,9 @@ async def test_compact_rejects_unsupported_configuration_and_shared_group_member ): context, manager = _compact_context( _compact_history(), - settings=settings, + provider_settings=provider_settings, + runner_type=runner_type, + compression_settings=compression_settings, unique_session=unique_session, ) event = FakeCompactEvent(group_id=group_id, role=role) @@ -352,29 +396,80 @@ async def test_compact_rejects_unsupported_configuration_and_shared_group_member provider_resolver.assert_not_awaited() +@pytest.mark.asyncio +async def test_compact_rejects_missing_active_conversation( + monkeypatch: pytest.MonkeyPatch, +): + context, manager = _compact_context(_compact_history()) + manager.get_curr_conversation_id.return_value = None + event = FakeCompactEvent() + provider_resolver = AsyncMock() + monkeypatch.setattr( + conversation_module, + "get_context_compression_provider", + provider_resolver, + ) + + await conversation_module.ConversationCommands(context).compact(event) + + assert "not in a conversation" in _result_text(event) + provider_resolver.assert_not_awaited() + manager.update_conversation.assert_not_awaited() + assert event.sent == [] + + +@pytest.mark.asyncio +async def test_compact_preserves_history_when_conversation_cannot_be_loaded( + monkeypatch: pytest.MonkeyPatch, +): + context, manager = _compact_context(_compact_history()) + manager.get_conversation.return_value = None + event = FakeCompactEvent() + provider = FakeCompressionProvider() + monkeypatch.setattr( + conversation_module, + "get_context_compression_provider", + AsyncMock(return_value=provider), + ) + + await conversation_module.ConversationCommands(context).compact(event) + + assert "could not be loaded" in _result_text(event) + manager.update_conversation.assert_not_awaited() + assert provider.call_count == 0 + assert [chain.type for chain in event.sent] == ["webchat_ephemeral"] + + @pytest.mark.asyncio @pytest.mark.parametrize( - ("platform_name", "sent_types"), - [("webchat", ["webchat_ephemeral", "agent_stats"]), ("telegram", [None])], + ("platform_name", "group_id", "role", "sent_types"), + [ + ("webchat", "", "member", ["webchat_ephemeral", "agent_stats"]), + ("telegram", "", "member", [None]), + ("webchat", "group-1", "admin", ["webchat_ephemeral", "agent_stats"]), + ], ) async def test_compact_writes_reduced_checkpoint_history_and_expected_stats( monkeypatch: pytest.MonkeyPatch, platform_name: str, + group_id: str, + role: str, sent_types: list[str | None], ): history = _compact_history() context, manager = _compact_context(history) - event = FakeCompactEvent(platform_name=platform_name) + event = FakeCompactEvent( + platform_name=platform_name, + group_id=group_id, + role=role, + ) provider = FakeCompressionProvider() - - async def get_provider(*args, **kwargs): - _ = args, kwargs - return provider + provider_resolver = AsyncMock(return_value=provider) monkeypatch.setattr( conversation_module, "get_context_compression_provider", - get_provider, + provider_resolver, ) await conversation_module.ConversationCommands(context).compact(event) @@ -387,6 +482,12 @@ async def get_provider(*args, **kwargs): assert {"role": "_checkpoint", "content": {"id": "cp-latest"}} in saved_history assert {"role": "_checkpoint", "content": {"id": "cp-old"}} not in saved_history assert provider.call_count == 1 + provider_resolver.assert_awaited_once_with( + "llm_compress", + "compression-provider", + context, + event, + ) assert [chain.type for chain in event.sent] == sent_types assert event.sent[0].get_plain_text() == "⏳ Compressing context..." if platform_name == "webchat": @@ -395,15 +496,19 @@ async def get_provider(*args, **kwargs): @pytest.mark.asyncio -@pytest.mark.parametrize("provider_failure", [False, True]) -async def test_compact_does_not_write_when_summary_fails_or_is_unchanged( +@pytest.mark.parametrize( + ("summary", "provider_failure"), + [("", False), ("", True), ("expanded summary " * 1000, False)], +) +async def test_compact_does_not_write_when_summary_fails_or_has_no_benefit( monkeypatch: pytest.MonkeyPatch, + summary: str, provider_failure: bool, ): history = _compact_history() context, manager = _compact_context(history) event = FakeCompactEvent() - provider = FakeCompressionProvider(summary="", fail=provider_failure) + provider = FakeCompressionProvider(summary=summary, fail=provider_failure) async def get_provider(*args, **kwargs): _ = args, kwargs @@ -423,14 +528,21 @@ async def get_provider(*args, **kwargs): @pytest.mark.asyncio -async def test_compact_reports_single_complete_round_as_not_enough_history( +@pytest.mark.parametrize( + "history", + [ + [], + [ + {"role": "user", "content": "question"}, + {"role": "assistant", "content": "answer"}, + {"role": "_checkpoint", "content": {"id": "cp-latest"}}, + ], + ], +) +async def test_compact_reports_insufficient_history_without_calling_provider( monkeypatch: pytest.MonkeyPatch, + history: list[dict], ): - history = [ - {"role": "user", "content": "question"}, - {"role": "assistant", "content": "answer"}, - {"role": "_checkpoint", "content": {"id": "cp-latest"}}, - ] context, manager = _compact_context(history) event = FakeCompactEvent() provider = FakeCompressionProvider() @@ -527,13 +639,53 @@ async def get_provider(*args, **kwargs): @pytest.mark.asyncio +async def test_compact_force_stop_after_storage_update_skips_stats_and_result( + monkeypatch: pytest.MonkeyPatch, +): + context, manager = _compact_context(_compact_history()) + event = FakeCompactEvent() + provider = FakeCompressionProvider() + + async def stop_after_update(*args, **kwargs): + _ = args, kwargs + event.stopped = True + + manager.update_conversation.side_effect = stop_after_update + monkeypatch.setattr( + conversation_module, + "get_context_compression_provider", + AsyncMock(return_value=provider), + ) + + await conversation_module.ConversationCommands(context).compact(event) + + manager.update_conversation.assert_awaited_once() + assert event.result is None + assert [chain.type for chain in event.sent] == ["webchat_ephemeral"] + + +@pytest.mark.asyncio +@pytest.mark.parametrize("stop_during_summary", [False, True]) async def test_compact_does_not_write_after_stop_request( monkeypatch: pytest.MonkeyPatch, + stop_during_summary: bool, ): context, manager = _compact_context(_compact_history()) - event = FakeCompactEvent(extras={"agent_stop_requested": True}) + event = FakeCompactEvent( + extras={} if stop_during_summary else {"agent_stop_requested": True} + ) provider = FakeCompressionProvider() + if stop_during_summary: + + async def request_stop_during_summary(**kwargs): + _ = kwargs + provider.call_count += 1 + event.extras["agent_stop_requested"] = True + return SimpleNamespace(completion_text="Concise summary.") + + provider.text_chat = request_stop_during_summary + async def get_provider(*args, **kwargs): _ = args, kwargs return provider @@ -547,7 +699,7 @@ async def get_provider(*args, **kwargs): await conversation_module.ConversationCommands(context).compact(event) manager.update_conversation.assert_not_awaited() - assert provider.call_count == 0 + assert provider.call_count == int(stop_during_summary) assert "cancelled" in _result_text(event) diff --git a/tests/unit/test_agent_runner_config.py b/tests/unit/test_agent_runner_config.py index 67b11e208e..0f6c62d0cc 100644 --- a/tests/unit/test_agent_runner_config.py +++ b/tests/unit/test_agent_runner_config.py @@ -32,7 +32,9 @@ def test_agent_runner_defaults_are_isolated_and_normalized(runner_type: str): "runner_type": runner_type, "config": second, } - if runner_type != "local": + if runner_type == "local": + assert second["compression"]["enable_manual_context_compression"] is False + else: assert "persona_id" not in second @@ -44,6 +46,7 @@ def test_switching_runner_type_discards_previous_runner_fields(): "provider_id": "legacy-provider", "persona_id": "legacy-persona", "model": {"provider_id": "chat-model"}, + "compression": {"enable_manual_context_compression": True}, "dify_api_key": "secret", "unexpected": True, }, @@ -57,6 +60,17 @@ def test_switching_runner_type_discards_previous_runner_fields(): assert "provider_id" not in normalized["config"] assert "persona_id" not in normalized["config"] assert "model" not in normalized["config"] + assert "compression" not in normalized["config"] + + switched_back = normalize_agent_runner( + {"runner_type": "local", "config": normalized["config"]} + ) + assert ( + switched_back["config"]["compression"][ + "enable_manual_context_compression" + ] + is False + ) @pytest.mark.asyncio @@ -88,9 +102,12 @@ async def test_agent_request_normalizes_incomplete_runner_config(): ) def test_each_runner_configuration_round_trips(tmp_path, runner_type: str): config = copy.deepcopy(DEFAULT_CONFIG) + runner_config = get_agent_runner_config_default(runner_type) + if runner_type == "local": + runner_config["compression"]["enable_manual_context_compression"] = True expected = { "runner_type": runner_type, - "config": get_agent_runner_config_default(runner_type), + "config": runner_config, } config["agent_runner"] = expected config_path = tmp_path / f"{runner_type}.json" @@ -120,6 +137,7 @@ def test_local_legacy_fields_are_fully_migrated(): "tool_call_timeout": 88, "sanitize_context_by_modalities": True, "context_limit_reached_strategy": "truncate_by_turns", + "enable_manual_context_compression": True, "llm_compress_instruction": "Summarize", "llm_compress_keep_recent_ratio": 0.2, "llm_compress_provider_id": "compressor", @@ -155,6 +173,7 @@ def test_local_legacy_fields_are_fully_migrated(): "max_turns": 20, "trim_turns": 3, "overflow_strategy": "truncate_by_turns", + "enable_manual_context_compression": True, "instruction": "Summarize", "keep_recent_ratio": 0.2, "provider_id": "compressor", @@ -177,6 +196,7 @@ def test_local_legacy_fields_are_fully_migrated(): "max_agent_step", "tool_call_timeout", "sanitize_context_by_modalities", + "enable_manual_context_compression", }.intersection(config["provider_settings"]) @@ -235,6 +255,7 @@ def test_local_migration_replaces_default_root_inserted_before_version_bump(): "max_turns": 24, "trim_turns": 4, "overflow_strategy": "llm_compress", + "enable_manual_context_compression": False, "instruction": "Keep decisions", "keep_recent_ratio": 0.15, "provider_id": "compressor", @@ -571,6 +592,7 @@ def test_new_agent_runner_config_is_authoritative_and_opaque_on_reload(tmp_path) config = copy.deepcopy(DEFAULT_CONFIG) config["config_version"] = 2 config["provider_settings"]["agent_runner_type"] = "coze" + config["provider_settings"]["enable_manual_context_compression"] = True config["agent_runner"] = { "runner_type": "dify", "config": { @@ -588,3 +610,4 @@ def test_new_agent_runner_config_is_authoritative_and_opaque_on_reload(tmp_path) assert loaded["agent_runner"]["config"]["dify_api_key"] == "saved-key" assert loaded["agent_runner"]["config"]["variables"] == {"nested": {"value": 1}} assert "agent_runner_type" not in loaded["provider_settings"] + assert "enable_manual_context_compression" not in loaded["provider_settings"] diff --git a/tests/unit/test_config.py b/tests/unit/test_config.py index 5c25c1ee40..061ff51f1f 100644 --- a/tests/unit/test_config.py +++ b/tests/unit/test_config.py @@ -9,7 +9,11 @@ import pytest from astrbot.core.config.astrbot_config import AstrBotConfig, RateLimitStrategy -from astrbot.core.config.default import DEFAULT_VALUE_MAP, get_local_permission_defaults +from astrbot.core.config.default import ( + CONFIG_METADATA_3, + DEFAULT_VALUE_MAP, + get_local_permission_defaults, +) from astrbot.core.config.i18n_utils import ConfigMetadataI18n from astrbot.core.utils.auth_password import ( DEFAULT_DASHBOARD_PASSWORD, @@ -994,6 +998,46 @@ def test_nested_object_schema(self, temp_config_path): class TestConfigMetadataI18n: """Tests for i18n utils.""" + @pytest.mark.parametrize("locale", ["en-US", "zh-CN", "ru-RU"]) + def test_manual_compression_metadata_uses_translated_runner_config_keys( + self, + locale: str, + ): + """Verify the embedded manual compression field and its translations.""" + field_key = ( + "agent_runner.config.compression.enable_manual_context_compression" + ) + metadata_item = CONFIG_METADATA_3["ai_group"]["metadata"][ + "truncate_and_compress" + ]["items"][field_key] + + assert metadata_item["type"] == "bool" + assert metadata_item["condition"] == { + "agent_runner.config.compression.overflow_strategy": "llm_compress", + "agent_runner.runner_type": "local", + } + + converted_item = ConfigMetadataI18n.convert_to_i18n_keys(CONFIG_METADATA_3)[ + "ai_group" + ]["metadata"]["truncate_and_compress"]["items"][field_key] + locale_path = ( + Path(__file__).parents[2] + / "dashboard" + / "src" + / "i18n" + / "locales" + / locale + / "features" + / "config-metadata.json" + ) + translations = json.loads(locale_path.read_text(encoding="utf-8")) + for attribute in ("description", "hint"): + translated_value = translations + for segment in converted_item[attribute].split("."): + translated_value = translated_value[segment] + assert isinstance(translated_value, str) + assert translated_value + def test_get_i18n_key(self): """Test generating i18n key.""" key = ConfigMetadataI18n._get_i18n_key( diff --git a/tests/unit/test_config_profile_service.py b/tests/unit/test_config_profile_service.py index bc1e2a2f3d..a34df469a1 100644 --- a/tests/unit/test_config_profile_service.py +++ b/tests/unit/test_config_profile_service.py @@ -4,6 +4,7 @@ import pytest +from astrbot.core.config.agent_runner import get_agent_runner_config_default from astrbot.dashboard.services.config_service import ConfigProfileService @@ -62,16 +63,22 @@ async def test_profile_mutations_await_config_manager() -> None: config_manager.create_conf.assert_not_awaited() lifecycle.reload_pipeline_scheduler.assert_not_awaited() - result = await service.create_profile( - "Profile", {"provider_settings": {"computer_use_runtime": "local"}} - ) + runner_config = get_agent_runner_config_default("local") + runner_config["compression"]["enable_manual_context_compression"] = True + profile_config = { + "timezone": "UTC", + "provider_settings": {"computer_use_runtime": "local"}, + "agent_runner": {"runner_type": "local", "config": runner_config}, + } + + result = await service.create_profile("Profile", profile_config) await service.rename_profile("profile-id", "Renamed") await service.delete_profile("profile-id") assert result == {"conf_id": "profile-id"} config_manager.create_conf.assert_awaited_once_with( name="Profile", - config={"provider_settings": {"computer_use_runtime": "local"}}, + config=profile_config, ) lifecycle.reload_pipeline_scheduler.assert_awaited_once_with("profile-id") config_manager.update_conf_info.assert_awaited_once_with( diff --git a/tests/unit/test_live_chat_service.py b/tests/unit/test_live_chat_service.py index 5cb3035689..05549f764a 100644 --- a/tests/unit/test_live_chat_service.py +++ b/tests/unit/test_live_chat_service.py @@ -287,6 +287,16 @@ async def send_json(payload: dict) -> None: "message_id": message_id, }, ) + await webchat_queue_mgr.put_back_queue( + message_id, + { + "type": "plain", + "data": '{"current_context_tokens": 42}', + "streaming": False, + "chain_type": "agent_stats", + "message_id": message_id, + }, + ) await webchat_queue_mgr.put_back_queue( message_id, { @@ -314,9 +324,17 @@ async def send_json(payload: dict) -> None: assert [ payload["data"] for payload in sent if payload.get("type") == "plain" ] == ["⏳ Compressing context...", "✅ Context compressed."] + assert [ + payload["data"] + for payload in sent + if payload.get("type") == "agent_stats" + ] == [{"current_context_tokens": 42}] service.save_bot_message.assert_awaited_once() - saved_parts = service.save_bot_message.await_args.args[1] - assert saved_parts == [{"type": "plain", "text": "✅ Context compressed."}] + save_args = service.save_bot_message.await_args.args + assert save_args[1] == [ + {"type": "plain", "text": "✅ Context compressed."} + ] + assert save_args[2] == {"current_context_tokens": 42} finally: if not task.done(): task.cancel() diff --git a/tests/unit/test_open_api_service_ws.py b/tests/unit/test_open_api_service_ws.py index adcee691cd..b6b89ddb00 100644 --- a/tests/unit/test_open_api_service_ws.py +++ b/tests/unit/test_open_api_service_ws.py @@ -1,7 +1,10 @@ +import asyncio from types import SimpleNamespace +from unittest.mock import AsyncMock import pytest +from astrbot.core.platform.sources.webchat.webchat_queue_mgr import webchat_queue_mgr from astrbot.dashboard.services.open_api_service import ( OpenApiService, OpenApiServiceError, @@ -147,6 +150,102 @@ async def close(_code: int, _reason: str) -> None: ] +@pytest.mark.asyncio +async def test_handle_chat_ws_send_forwards_ephemeral_and_persists_terminal_with_stats(): + service = _service() + session_id = "compact-openapi-session" + message_id = "compact-openapi-request" + sent: list[dict] = [] + errors: list[tuple[str, str]] = [] + bridge = _bridge() + bridge.build_user_message_parts = AsyncMock( + return_value=[{"type": "plain", "text": "hello"}] + ) + bridge.save_bot_message = AsyncMock(return_value=None) + service.prepare_chat_send = AsyncMock( + return_value=("alice", session_id, None) + ) + service.update_session_config_route = AsyncMock(return_value=None) + + async def send_json(payload: dict) -> None: + sent.append(payload) + + async def send_error(message: str, code: str) -> None: + errors.append((message, code)) + + input_queue = webchat_queue_mgr.get_or_create_queue(session_id) + task = asyncio.create_task( + service.handle_chat_ws_send( + post_data={ + "message": "hello", + "session_id": session_id, + "message_id": message_id, + }, + conf_list=[], + chat_bridge=bridge, + send_json=send_json, + send_error=send_error, + ) + ) + + try: + await asyncio.wait_for(input_queue.get(), timeout=1) + for payload in ( + { + "type": "plain", + "data": "⏳ Compressing context...", + "streaming": False, + "chain_type": "webchat_ephemeral", + "message_id": message_id, + }, + { + "type": "plain", + "data": '{"current_context_tokens": 42}', + "streaming": False, + "chain_type": "agent_stats", + "message_id": message_id, + }, + { + "type": "plain", + "data": "✅ Context compressed.", + "streaming": False, + "message_id": message_id, + }, + { + "type": "end", + "data": "", + "streaming": False, + "message_id": message_id, + }, + ): + assert await webchat_queue_mgr.put_back_queue(message_id, payload) + + await asyncio.wait_for(task, timeout=1) + + assert errors == [] + assert [ + payload["data"] for payload in sent if payload.get("type") == "plain" + ] == ["⏳ Compressing context...", "✅ Context compressed."] + assert [ + payload["data"] + for payload in sent + if payload.get("type") == "agent_stats" + ] == [{"current_context_tokens": 42}] + bridge.save_bot_message.assert_awaited_once() + save_args = bridge.save_bot_message.await_args.args + assert save_args[0] == session_id + assert save_args[1] == [ + {"type": "plain", "text": "✅ Context compressed."} + ] + assert save_args[2] == {"current_context_tokens": 42} + assert save_args[3] == {} + finally: + if not task.done(): + task.cancel() + await asyncio.gather(task, return_exceptions=True) + webchat_queue_mgr.remove_queues(session_id) + + @pytest.mark.asyncio async def test_prepare_chat_send_rejects_configured_admin_username(): """The shared HTTP/WS boundary must reject administrator impersonation.""" From 1c79dab2b39cfc23ba024645e1607b62a7aadb89 Mon Sep 17 00:00:00 2001 From: C10H14N2O5 <100066858+C10H14N2O5@users.noreply.github.com> Date: Sun, 30 Aug 2026 19:56:32 +0800 Subject: [PATCH 07/16] fix: preserve legacy runner config migration --- astrbot/core/utils/migra_helper.py | 11 ++++++++++- tests/unit/test_agent_runner_config.py | 4 +++- 2 files changed, 13 insertions(+), 2 deletions(-) diff --git a/astrbot/core/utils/migra_helper.py b/astrbot/core/utils/migra_helper.py index 060dac10af..1a9b5c0779 100644 --- a/astrbot/core/utils/migra_helper.py +++ b/astrbot/core/utils/migra_helper.py @@ -170,9 +170,18 @@ def _migrate_agent_runner_config( "runner_type": "local", "config": get_agent_runner_config_default("local"), } + comparable_agent_runner = copy.deepcopy(existing_agent_runner) + if isinstance(comparable_agent_runner, dict): + comparable_runner_config = comparable_agent_runner.get("config") + if isinstance(comparable_runner_config, dict): + comparable_compression = comparable_runner_config.get("compression") + if isinstance(comparable_compression, dict): + comparable_compression.setdefault( + "enable_manual_context_compression", False + ) default_root_inserted_before_migration = ( legacy_version - and existing_agent_runner == default_local_agent_runner + and comparable_agent_runner == default_local_agent_runner and any(key in provider_settings for key in _LEGACY_AGENT_RUNNER_SETTING_KEYS) ) diff --git a/tests/unit/test_agent_runner_config.py b/tests/unit/test_agent_runner_config.py index 0f6c62d0cc..319ed18303 100644 --- a/tests/unit/test_agent_runner_config.py +++ b/tests/unit/test_agent_runner_config.py @@ -201,6 +201,8 @@ def test_local_legacy_fields_are_fully_migrated(): def test_local_migration_replaces_default_root_inserted_before_version_bump(): + legacy_default_config = get_agent_runner_config_default("local") + legacy_default_config["compression"].pop("enable_manual_context_compression") config = { "config_version": 2, "provider": [ @@ -225,7 +227,7 @@ def test_local_migration_replaces_default_root_inserted_before_version_bump(): }, "agent_runner": { "runner_type": "local", - "config": get_agent_runner_config_default("local"), + "config": legacy_default_config, }, } From 6b40ff85b889d4582585d5da5578539775c4670b Mon Sep 17 00:00:00 2001 From: C10H14N2O5 <100066858+C10H14N2O5@users.noreply.github.com> Date: Sun, 30 Aug 2026 19:56:47 +0800 Subject: [PATCH 08/16] fix: preserve context token usage compatibility --- astrbot/core/agent/context/manager.py | 9 +++++---- astrbot/core/agent/context/token_counter.py | 10 +++++----- astrbot/core/agent/runners/tool_loop_agent_runner.py | 2 +- tests/agent/test_context_manager.py | 11 +++++++---- tests/agent/test_token_counter.py | 6 +++--- tests/test_tool_loop_agent_runner.py | 4 ++-- 6 files changed, 23 insertions(+), 19 deletions(-) diff --git a/astrbot/core/agent/context/manager.py b/astrbot/core/agent/context/manager.py index d8e29fc900..540806e453 100644 --- a/astrbot/core/agent/context/manager.py +++ b/astrbot/core/agent/context/manager.py @@ -46,14 +46,14 @@ def __init__( async def process( self, messages: list[Message], - reported_token_usage: int = 0, + trusted_token_usage: int = 0, force_compress: bool = False, ) -> list[Message]: """Process the messages. Args: messages: The original message list. - reported_token_usage: Token usage reported by the previous provider call. + trusted_token_usage: Token usage reported by the previous provider call. force_compress: Whether to bypass automatic limits and run the configured compressor immediately without a truncation fallback. @@ -82,13 +82,14 @@ async def process( # 2. 基于 token 的压缩 if self.config.max_context_tokens > 0: total_tokens = self.token_counter.count_tokens( - result, reported_token_usage + result, trusted_token_usage ) if self.compressor.should_compress( result, total_tokens, self.config.max_context_tokens ): - result = await self._run_compression(result, total_tokens) + estimated_tokens = self.token_counter.count_tokens(result) + result = await self._run_compression(result, estimated_tokens) return result except Exception as e: diff --git a/astrbot/core/agent/context/token_counter.py b/astrbot/core/agent/context/token_counter.py index 696e219959..31c274af64 100644 --- a/astrbot/core/agent/context/token_counter.py +++ b/astrbot/core/agent/context/token_counter.py @@ -13,13 +13,13 @@ class TokenCounter(Protocol): """ def count_tokens( - self, messages: list[Message], reported_token_usage: int = 0 + self, messages: list[Message], trusted_token_usage: int = 0 ) -> int: """Count the total tokens in the message list. Args: messages: The message list. - reported_token_usage: The total token usage that LLM API returned. + trusted_token_usage: The total token usage that LLM API returned. For some cases, this value is more accurate. But some API does not return it, so the value defaults to 0. @@ -54,10 +54,10 @@ class EstimateTokenCounter: """ def count_tokens( - self, messages: list[Message], reported_token_usage: int = 0 + self, messages: list[Message], trusted_token_usage: int = 0 ) -> int: - if reported_token_usage > 0: - return reported_token_usage + if trusted_token_usage > 0: + return trusted_token_usage total = 0 for msg in messages: diff --git a/astrbot/core/agent/runners/tool_loop_agent_runner.py b/astrbot/core/agent/runners/tool_loop_agent_runner.py index e473d4af54..ed5bb3df62 100644 --- a/astrbot/core/agent/runners/tool_loop_agent_runner.py +++ b/astrbot/core/agent/runners/tool_loop_agent_runner.py @@ -829,7 +829,7 @@ async def step(self): processed_messages = await self._await_or_stop( self.request_context_manager.process( self.run_context.messages, - reported_token_usage=token_usage, + trusted_token_usage=token_usage, ) ) if processed_messages is None: diff --git a/tests/agent/test_context_manager.py b/tests/agent/test_context_manager.py index 2ccfb51555..2ed971931d 100644 --- a/tests/agent/test_context_manager.py +++ b/tests/agent/test_context_manager.py @@ -635,7 +635,8 @@ def mock_should_compress(*args, **kwargs): assert len(result) <= len(messages) @pytest.mark.asyncio - async def test_reported_usage_triggers_compression_before_provider_call(self): + async def test_trusted_usage_triggers_compression_before_provider_call(self, caplog): + caplog.set_level("INFO", logger="astrbot") config = ContextConfig(max_context_tokens=100, truncate_turns=1) manager = ContextManager(config) messages = [self.create_message("user", "short")] @@ -644,16 +645,18 @@ async def test_reported_usage_triggers_compression_before_provider_call(self): mock_compressor.should_compress = MagicMock(side_effect=[True, False]) manager.compressor = mock_compressor - result = await manager.process(messages, reported_token_usage=83) + result = await manager.process(messages, trusted_token_usage=83) first_check = mock_compressor.should_compress.call_args_list[0] assert first_check.args == (messages, 83, 100) mock_compressor.assert_awaited_once_with(messages) assert result == compressed + assert "Compress completed." in caplog.text + assert " 83 ->" not in caplog.text @pytest.mark.asyncio async def test_force_compression_bypasses_automatic_guards(self): - """Forced compression ignores limits, reported usage, and truncation.""" + """Forced compression ignores limits, trusted usage, and truncation.""" config = ContextConfig(max_context_tokens=0, enforce_max_turns=1) manager = ContextManager(config) messages = self.create_messages(6) @@ -669,7 +672,7 @@ async def test_force_compression_bypasses_automatic_guards(self): ): result = await manager.process( messages, - reported_token_usage=999, + trusted_token_usage=999, force_compress=True, ) diff --git a/tests/agent/test_token_counter.py b/tests/agent/test_token_counter.py index 6375b45452..0b72403e1b 100644 --- a/tests/agent/test_token_counter.py +++ b/tests/agent/test_token_counter.py @@ -127,8 +127,8 @@ def test_plain_text_estimate_unchanged(self): assert counter.count_tokens([_msg("user", "你" * 100)]) == 60 -class TestReportedUsage: - def test_reported_overrides(self): +class TestTrustedUsage: + def test_trusted_overrides(self): """如果 API 返回了 token 数,直接用它不做估算。""" msg = _msg( "user", @@ -139,7 +139,7 @@ def test_reported_overrides(self): ), ], ) - tokens = counter.count_tokens([msg], reported_token_usage=42) + tokens = counter.count_tokens([msg], trusted_token_usage=42) assert tokens == 42 diff --git a/tests/test_tool_loop_agent_runner.py b/tests/test_tool_loop_agent_runner.py index 03e65bc928..ef8e818bb5 100644 --- a/tests/test_tool_loop_agent_runner.py +++ b/tests/test_tool_loop_agent_runner.py @@ -573,7 +573,7 @@ async def test_max_step_final_request_includes_limit_prompt( streaming=False, ) - async def snapshot_context_manager(messages, reported_token_usage=0): + async def snapshot_context_manager(messages, trusted_token_usage=0): return list(messages) runner.request_context_manager.process = snapshot_context_manager @@ -603,7 +603,7 @@ async def test_tool_loop_next_request_includes_tool_result( streaming=False, ) - async def snapshot_context_manager(messages, reported_token_usage=0): + async def snapshot_context_manager(messages, trusted_token_usage=0): return list(messages) runner.request_context_manager.process = snapshot_context_manager From ffaa20c89b592239e69ca9d2ebd04c4aae7be18a Mon Sep 17 00:00:00 2001 From: C10H14N2O5 <100066858+C10H14N2O5@users.noreply.github.com> Date: Sun, 30 Aug 2026 20:22:18 +0800 Subject: [PATCH 09/16] fix: retain token counter keyword compatibility --- astrbot/core/agent/context/token_counter.py | 19 ++++++++++++++++--- tests/agent/test_token_counter.py | 5 +++++ 2 files changed, 21 insertions(+), 3 deletions(-) diff --git a/astrbot/core/agent/context/token_counter.py b/astrbot/core/agent/context/token_counter.py index 31c274af64..5e4839df27 100644 --- a/astrbot/core/agent/context/token_counter.py +++ b/astrbot/core/agent/context/token_counter.py @@ -54,10 +54,23 @@ class EstimateTokenCounter: """ def count_tokens( - self, messages: list[Message], trusted_token_usage: int = 0 + self, + messages: list[Message], + reported_token_usage: int = 0, + **legacy_usage: int, ) -> int: - if trusted_token_usage > 0: - return trusted_token_usage + unexpected = legacy_usage.keys() - {"trusted_token_usage"} + if unexpected: + name = next(iter(unexpected)) + raise TypeError( + f"count_tokens() got an unexpected keyword argument '{name}'" + ) + if reported_token_usage <= 0: + reported_token_usage = legacy_usage.get( + "trusted_token_usage", reported_token_usage + ) + if reported_token_usage > 0: + return reported_token_usage total = 0 for msg in messages: diff --git a/tests/agent/test_token_counter.py b/tests/agent/test_token_counter.py index 0b72403e1b..555dc31807 100644 --- a/tests/agent/test_token_counter.py +++ b/tests/agent/test_token_counter.py @@ -141,6 +141,11 @@ def test_trusted_overrides(self): ) tokens = counter.count_tokens([msg], trusted_token_usage=42) assert tokens == 42 + assert counter.count_tokens([msg], reported_token_usage=43) == 43 + assert ( + counter.count_tokens([msg], reported_token_usage=43, trusted_token_usage=0) + == 43 + ) class TestToolCalls: From bbcd6918f5dab0c3b435da4a41a36219ebbc1b79 Mon Sep 17 00:00:00 2001 From: C10H14N2O5 <100066858+C10H14N2O5@users.noreply.github.com> Date: Sun, 30 Aug 2026 20:51:23 +0800 Subject: [PATCH 10/16] fix: restore CodeQL-safe token usage naming --- astrbot/core/agent/context/manager.py | 6 +++--- astrbot/core/agent/context/token_counter.py | 19 +++---------------- .../agent/runners/tool_loop_agent_runner.py | 2 +- tests/agent/test_context_manager.py | 10 ++++++---- tests/agent/test_token_counter.py | 11 +++-------- tests/test_tool_loop_agent_runner.py | 4 ++-- 6 files changed, 18 insertions(+), 34 deletions(-) diff --git a/astrbot/core/agent/context/manager.py b/astrbot/core/agent/context/manager.py index 540806e453..9311aa3fee 100644 --- a/astrbot/core/agent/context/manager.py +++ b/astrbot/core/agent/context/manager.py @@ -46,14 +46,14 @@ def __init__( async def process( self, messages: list[Message], - trusted_token_usage: int = 0, + reported_token_usage: int = 0, force_compress: bool = False, ) -> list[Message]: """Process the messages. Args: messages: The original message list. - trusted_token_usage: Token usage reported by the previous provider call. + reported_token_usage: Token usage reported by the previous provider call. force_compress: Whether to bypass automatic limits and run the configured compressor immediately without a truncation fallback. @@ -82,7 +82,7 @@ async def process( # 2. 基于 token 的压缩 if self.config.max_context_tokens > 0: total_tokens = self.token_counter.count_tokens( - result, trusted_token_usage + result, reported_token_usage ) if self.compressor.should_compress( diff --git a/astrbot/core/agent/context/token_counter.py b/astrbot/core/agent/context/token_counter.py index 5e4839df27..696e219959 100644 --- a/astrbot/core/agent/context/token_counter.py +++ b/astrbot/core/agent/context/token_counter.py @@ -13,13 +13,13 @@ class TokenCounter(Protocol): """ def count_tokens( - self, messages: list[Message], trusted_token_usage: int = 0 + self, messages: list[Message], reported_token_usage: int = 0 ) -> int: """Count the total tokens in the message list. Args: messages: The message list. - trusted_token_usage: The total token usage that LLM API returned. + reported_token_usage: The total token usage that LLM API returned. For some cases, this value is more accurate. But some API does not return it, so the value defaults to 0. @@ -54,21 +54,8 @@ class EstimateTokenCounter: """ def count_tokens( - self, - messages: list[Message], - reported_token_usage: int = 0, - **legacy_usage: int, + self, messages: list[Message], reported_token_usage: int = 0 ) -> int: - unexpected = legacy_usage.keys() - {"trusted_token_usage"} - if unexpected: - name = next(iter(unexpected)) - raise TypeError( - f"count_tokens() got an unexpected keyword argument '{name}'" - ) - if reported_token_usage <= 0: - reported_token_usage = legacy_usage.get( - "trusted_token_usage", reported_token_usage - ) if reported_token_usage > 0: return reported_token_usage diff --git a/astrbot/core/agent/runners/tool_loop_agent_runner.py b/astrbot/core/agent/runners/tool_loop_agent_runner.py index ed5bb3df62..e473d4af54 100644 --- a/astrbot/core/agent/runners/tool_loop_agent_runner.py +++ b/astrbot/core/agent/runners/tool_loop_agent_runner.py @@ -829,7 +829,7 @@ async def step(self): processed_messages = await self._await_or_stop( self.request_context_manager.process( self.run_context.messages, - trusted_token_usage=token_usage, + reported_token_usage=token_usage, ) ) if processed_messages is None: diff --git a/tests/agent/test_context_manager.py b/tests/agent/test_context_manager.py index 2ed971931d..0825be53e5 100644 --- a/tests/agent/test_context_manager.py +++ b/tests/agent/test_context_manager.py @@ -635,7 +635,9 @@ def mock_should_compress(*args, **kwargs): assert len(result) <= len(messages) @pytest.mark.asyncio - async def test_trusted_usage_triggers_compression_before_provider_call(self, caplog): + async def test_reported_usage_triggers_compression_before_provider_call( + self, caplog + ): caplog.set_level("INFO", logger="astrbot") config = ContextConfig(max_context_tokens=100, truncate_turns=1) manager = ContextManager(config) @@ -645,7 +647,7 @@ async def test_trusted_usage_triggers_compression_before_provider_call(self, cap mock_compressor.should_compress = MagicMock(side_effect=[True, False]) manager.compressor = mock_compressor - result = await manager.process(messages, trusted_token_usage=83) + result = await manager.process(messages, reported_token_usage=83) first_check = mock_compressor.should_compress.call_args_list[0] assert first_check.args == (messages, 83, 100) @@ -656,7 +658,7 @@ async def test_trusted_usage_triggers_compression_before_provider_call(self, cap @pytest.mark.asyncio async def test_force_compression_bypasses_automatic_guards(self): - """Forced compression ignores limits, trusted usage, and truncation.""" + """Forced compression ignores limits, reported usage, and truncation.""" config = ContextConfig(max_context_tokens=0, enforce_max_turns=1) manager = ContextManager(config) messages = self.create_messages(6) @@ -672,7 +674,7 @@ async def test_force_compression_bypasses_automatic_guards(self): ): result = await manager.process( messages, - trusted_token_usage=999, + reported_token_usage=999, force_compress=True, ) diff --git a/tests/agent/test_token_counter.py b/tests/agent/test_token_counter.py index 555dc31807..6375b45452 100644 --- a/tests/agent/test_token_counter.py +++ b/tests/agent/test_token_counter.py @@ -127,8 +127,8 @@ def test_plain_text_estimate_unchanged(self): assert counter.count_tokens([_msg("user", "你" * 100)]) == 60 -class TestTrustedUsage: - def test_trusted_overrides(self): +class TestReportedUsage: + def test_reported_overrides(self): """如果 API 返回了 token 数,直接用它不做估算。""" msg = _msg( "user", @@ -139,13 +139,8 @@ def test_trusted_overrides(self): ), ], ) - tokens = counter.count_tokens([msg], trusted_token_usage=42) + tokens = counter.count_tokens([msg], reported_token_usage=42) assert tokens == 42 - assert counter.count_tokens([msg], reported_token_usage=43) == 43 - assert ( - counter.count_tokens([msg], reported_token_usage=43, trusted_token_usage=0) - == 43 - ) class TestToolCalls: diff --git a/tests/test_tool_loop_agent_runner.py b/tests/test_tool_loop_agent_runner.py index ef8e818bb5..03e65bc928 100644 --- a/tests/test_tool_loop_agent_runner.py +++ b/tests/test_tool_loop_agent_runner.py @@ -573,7 +573,7 @@ async def test_max_step_final_request_includes_limit_prompt( streaming=False, ) - async def snapshot_context_manager(messages, trusted_token_usage=0): + async def snapshot_context_manager(messages, reported_token_usage=0): return list(messages) runner.request_context_manager.process = snapshot_context_manager @@ -603,7 +603,7 @@ async def test_tool_loop_next_request_includes_tool_result( streaming=False, ) - async def snapshot_context_manager(messages, trusted_token_usage=0): + async def snapshot_context_manager(messages, reported_token_usage=0): return list(messages) runner.request_context_manager.process = snapshot_context_manager From c1f3696da8536b5fbbe2b006c2ff8eba79e679b5 Mon Sep 17 00:00:00 2001 From: C10H14N2O5 <100066858+C10H14N2O5@users.noreply.github.com> Date: Sun, 30 Aug 2026 21:21:14 +0800 Subject: [PATCH 11/16] fix: preserve trusted token usage interface --- astrbot/core/agent/context/manager.py | 21 +++++++++++-------- astrbot/core/agent/context/token_counter.py | 10 ++++----- .../agent/runners/tool_loop_agent_runner.py | 2 +- tests/agent/test_context_manager.py | 8 +++---- tests/agent/test_token_counter.py | 10 ++++++--- tests/test_tool_loop_agent_runner.py | 4 ++-- 6 files changed, 31 insertions(+), 24 deletions(-) diff --git a/astrbot/core/agent/context/manager.py b/astrbot/core/agent/context/manager.py index 9311aa3fee..c9c548fbf6 100644 --- a/astrbot/core/agent/context/manager.py +++ b/astrbot/core/agent/context/manager.py @@ -46,14 +46,14 @@ def __init__( async def process( self, messages: list[Message], - reported_token_usage: int = 0, + trusted_token_usage: int = 0, force_compress: bool = False, ) -> list[Message]: """Process the messages. Args: messages: The original message list. - reported_token_usage: Token usage reported by the previous provider call. + trusted_token_usage: Token usage reported by the previous provider call. force_compress: Whether to bypass automatic limits and run the configured compressor immediately without a truncation fallback. @@ -72,24 +72,27 @@ async def process( ) if force_compress: - total_tokens = self.token_counter.count_tokens(result) + estimated_context_tokens = self.token_counter.count_tokens(result) return await self._run_compression( result, - total_tokens, + estimated_context_tokens, allow_halving_fallback=False, ) # 2. 基于 token 的压缩 if self.config.max_context_tokens > 0: - total_tokens = self.token_counter.count_tokens( - result, reported_token_usage + threshold_tokens = self.token_counter.count_tokens( + result, trusted_token_usage ) if self.compressor.should_compress( - result, total_tokens, self.config.max_context_tokens + result, threshold_tokens, self.config.max_context_tokens ): - estimated_tokens = self.token_counter.count_tokens(result) - result = await self._run_compression(result, estimated_tokens) + estimated_context_tokens = self.token_counter.count_tokens(result) + result = await self._run_compression( + result, + estimated_context_tokens, + ) return result except Exception as e: diff --git a/astrbot/core/agent/context/token_counter.py b/astrbot/core/agent/context/token_counter.py index 696e219959..31c274af64 100644 --- a/astrbot/core/agent/context/token_counter.py +++ b/astrbot/core/agent/context/token_counter.py @@ -13,13 +13,13 @@ class TokenCounter(Protocol): """ def count_tokens( - self, messages: list[Message], reported_token_usage: int = 0 + self, messages: list[Message], trusted_token_usage: int = 0 ) -> int: """Count the total tokens in the message list. Args: messages: The message list. - reported_token_usage: The total token usage that LLM API returned. + trusted_token_usage: The total token usage that LLM API returned. For some cases, this value is more accurate. But some API does not return it, so the value defaults to 0. @@ -54,10 +54,10 @@ class EstimateTokenCounter: """ def count_tokens( - self, messages: list[Message], reported_token_usage: int = 0 + self, messages: list[Message], trusted_token_usage: int = 0 ) -> int: - if reported_token_usage > 0: - return reported_token_usage + if trusted_token_usage > 0: + return trusted_token_usage total = 0 for msg in messages: diff --git a/astrbot/core/agent/runners/tool_loop_agent_runner.py b/astrbot/core/agent/runners/tool_loop_agent_runner.py index e473d4af54..ed5bb3df62 100644 --- a/astrbot/core/agent/runners/tool_loop_agent_runner.py +++ b/astrbot/core/agent/runners/tool_loop_agent_runner.py @@ -829,7 +829,7 @@ async def step(self): processed_messages = await self._await_or_stop( self.request_context_manager.process( self.run_context.messages, - reported_token_usage=token_usage, + trusted_token_usage=token_usage, ) ) if processed_messages is None: diff --git a/tests/agent/test_context_manager.py b/tests/agent/test_context_manager.py index 0825be53e5..aa46158acf 100644 --- a/tests/agent/test_context_manager.py +++ b/tests/agent/test_context_manager.py @@ -635,7 +635,7 @@ def mock_should_compress(*args, **kwargs): assert len(result) <= len(messages) @pytest.mark.asyncio - async def test_reported_usage_triggers_compression_before_provider_call( + async def test_trusted_usage_triggers_compression_before_provider_call( self, caplog ): caplog.set_level("INFO", logger="astrbot") @@ -647,7 +647,7 @@ async def test_reported_usage_triggers_compression_before_provider_call( mock_compressor.should_compress = MagicMock(side_effect=[True, False]) manager.compressor = mock_compressor - result = await manager.process(messages, reported_token_usage=83) + result = await manager.process(messages, trusted_token_usage=83) first_check = mock_compressor.should_compress.call_args_list[0] assert first_check.args == (messages, 83, 100) @@ -658,7 +658,7 @@ async def test_reported_usage_triggers_compression_before_provider_call( @pytest.mark.asyncio async def test_force_compression_bypasses_automatic_guards(self): - """Forced compression ignores limits, reported usage, and truncation.""" + """Forced compression ignores limits, trusted usage, and truncation.""" config = ContextConfig(max_context_tokens=0, enforce_max_turns=1) manager = ContextManager(config) messages = self.create_messages(6) @@ -674,7 +674,7 @@ async def test_force_compression_bypasses_automatic_guards(self): ): result = await manager.process( messages, - reported_token_usage=999, + trusted_token_usage=999, force_compress=True, ) diff --git a/tests/agent/test_token_counter.py b/tests/agent/test_token_counter.py index 6375b45452..736b2cb888 100644 --- a/tests/agent/test_token_counter.py +++ b/tests/agent/test_token_counter.py @@ -1,5 +1,7 @@ """Tests for EstimateTokenCounter multimodal support.""" +import pytest + from astrbot.core.agent.context.token_counter import ( AUDIO_TOKEN_ESTIMATE, IMAGE_TOKEN_ESTIMATE, @@ -127,8 +129,8 @@ def test_plain_text_estimate_unchanged(self): assert counter.count_tokens([_msg("user", "你" * 100)]) == 60 -class TestReportedUsage: - def test_reported_overrides(self): +class TestTrustedUsage: + def test_trusted_overrides(self): """如果 API 返回了 token 数,直接用它不做估算。""" msg = _msg( "user", @@ -139,8 +141,10 @@ def test_reported_overrides(self): ), ], ) - tokens = counter.count_tokens([msg], reported_token_usage=42) + tokens = counter.count_tokens([msg], trusted_token_usage=42) assert tokens == 42 + with pytest.raises(TypeError, match="reported_token_usage"): + counter.count_tokens([msg], reported_token_usage=43) class TestToolCalls: diff --git a/tests/test_tool_loop_agent_runner.py b/tests/test_tool_loop_agent_runner.py index 03e65bc928..ef8e818bb5 100644 --- a/tests/test_tool_loop_agent_runner.py +++ b/tests/test_tool_loop_agent_runner.py @@ -573,7 +573,7 @@ async def test_max_step_final_request_includes_limit_prompt( streaming=False, ) - async def snapshot_context_manager(messages, reported_token_usage=0): + async def snapshot_context_manager(messages, trusted_token_usage=0): return list(messages) runner.request_context_manager.process = snapshot_context_manager @@ -603,7 +603,7 @@ async def test_tool_loop_next_request_includes_tool_result( streaming=False, ) - async def snapshot_context_manager(messages, reported_token_usage=0): + async def snapshot_context_manager(messages, trusted_token_usage=0): return list(messages) runner.request_context_manager.process = snapshot_context_manager From 97323a91437c473677deb2b30469518b45135548 Mon Sep 17 00:00:00 2001 From: C10H14N2O5 <100066858+C10H14N2O5@users.noreply.github.com> Date: Sun, 30 Aug 2026 21:27:35 +0800 Subject: [PATCH 12/16] fix: isolate reported token usage --- astrbot/core/agent/context/manager.py | 3 ++- astrbot/core/agent/context/token_counter.py | 10 +++++----- tests/agent/test_token_counter.py | 10 +++------- 3 files changed, 10 insertions(+), 13 deletions(-) diff --git a/astrbot/core/agent/context/manager.py b/astrbot/core/agent/context/manager.py index c9c548fbf6..cb96fdf88c 100644 --- a/astrbot/core/agent/context/manager.py +++ b/astrbot/core/agent/context/manager.py @@ -82,7 +82,8 @@ async def process( # 2. 基于 token 的压缩 if self.config.max_context_tokens > 0: threshold_tokens = self.token_counter.count_tokens( - result, trusted_token_usage + result, + reported_token_usage=trusted_token_usage, ) if self.compressor.should_compress( diff --git a/astrbot/core/agent/context/token_counter.py b/astrbot/core/agent/context/token_counter.py index 31c274af64..696e219959 100644 --- a/astrbot/core/agent/context/token_counter.py +++ b/astrbot/core/agent/context/token_counter.py @@ -13,13 +13,13 @@ class TokenCounter(Protocol): """ def count_tokens( - self, messages: list[Message], trusted_token_usage: int = 0 + self, messages: list[Message], reported_token_usage: int = 0 ) -> int: """Count the total tokens in the message list. Args: messages: The message list. - trusted_token_usage: The total token usage that LLM API returned. + reported_token_usage: The total token usage that LLM API returned. For some cases, this value is more accurate. But some API does not return it, so the value defaults to 0. @@ -54,10 +54,10 @@ class EstimateTokenCounter: """ def count_tokens( - self, messages: list[Message], trusted_token_usage: int = 0 + self, messages: list[Message], reported_token_usage: int = 0 ) -> int: - if trusted_token_usage > 0: - return trusted_token_usage + if reported_token_usage > 0: + return reported_token_usage total = 0 for msg in messages: diff --git a/tests/agent/test_token_counter.py b/tests/agent/test_token_counter.py index 736b2cb888..6375b45452 100644 --- a/tests/agent/test_token_counter.py +++ b/tests/agent/test_token_counter.py @@ -1,7 +1,5 @@ """Tests for EstimateTokenCounter multimodal support.""" -import pytest - from astrbot.core.agent.context.token_counter import ( AUDIO_TOKEN_ESTIMATE, IMAGE_TOKEN_ESTIMATE, @@ -129,8 +127,8 @@ def test_plain_text_estimate_unchanged(self): assert counter.count_tokens([_msg("user", "你" * 100)]) == 60 -class TestTrustedUsage: - def test_trusted_overrides(self): +class TestReportedUsage: + def test_reported_overrides(self): """如果 API 返回了 token 数,直接用它不做估算。""" msg = _msg( "user", @@ -141,10 +139,8 @@ def test_trusted_overrides(self): ), ], ) - tokens = counter.count_tokens([msg], trusted_token_usage=42) + tokens = counter.count_tokens([msg], reported_token_usage=42) assert tokens == 42 - with pytest.raises(TypeError, match="reported_token_usage"): - counter.count_tokens([msg], reported_token_usage=43) class TestToolCalls: From f2d59c9231fb611b3364cd3c5d34f41573a2f15c Mon Sep 17 00:00:00 2001 From: C10H14N2O5 <100066858+C10H14N2O5@users.noreply.github.com> Date: Sun, 30 Aug 2026 21:39:23 +0800 Subject: [PATCH 13/16] chore: rerun CI From 95c325cd65fb3115a7540cbf1416bf5b876eb15d Mon Sep 17 00:00:00 2001 From: C10H14N2O5 <100066858+C10H14N2O5@users.noreply.github.com> Date: Sun, 20 Sep 2026 00:21:12 +0800 Subject: [PATCH 14/16] fix: adapt manual compression to current upstream --- astrbot/core/utils/migra_helper.py | 11 ++++++- .../ja-JP/features/config-metadata.json | 4 +++ tests/unit/test_agent_runner_config.py | 32 +++++++++++++++++-- tests/unit/test_config.py | 2 +- 4 files changed, 45 insertions(+), 4 deletions(-) diff --git a/astrbot/core/utils/migra_helper.py b/astrbot/core/utils/migra_helper.py index 1a9b5c0779..2c40419c69 100644 --- a/astrbot/core/utils/migra_helper.py +++ b/astrbot/core/utils/migra_helper.py @@ -171,7 +171,7 @@ def _migrate_agent_runner_config( "config": get_agent_runner_config_default("local"), } comparable_agent_runner = copy.deepcopy(existing_agent_runner) - if isinstance(comparable_agent_runner, dict): + if legacy_version and isinstance(comparable_agent_runner, dict): comparable_runner_config = comparable_agent_runner.get("config") if isinstance(comparable_runner_config, dict): comparable_compression = comparable_runner_config.get("compression") @@ -179,6 +179,15 @@ def _migrate_agent_runner_config( comparable_compression.setdefault( "enable_manual_context_compression", False ) + comparable_misc = comparable_runner_config.get("misc") + if ( + isinstance(comparable_misc, dict) + and comparable_misc.get("max_steps") == 30 + ): + # Recognize defaults inserted before the v4 step-limit upgrade. + comparable_misc["max_steps"] = default_local_agent_runner["config"][ + "misc" + ]["max_steps"] default_root_inserted_before_migration = ( legacy_version and comparable_agent_runner == default_local_agent_runner diff --git a/dashboard/src/i18n/locales/ja-JP/features/config-metadata.json b/dashboard/src/i18n/locales/ja-JP/features/config-metadata.json index 8c44096b6d..2aea5398a6 100644 --- a/dashboard/src/i18n/locales/ja-JP/features/config-metadata.json +++ b/dashboard/src/i18n/locales/ja-JP/features/config-metadata.json @@ -468,6 +468,10 @@ ], "hint": "通常の会話履歴では、「圧縮前に保持する最大会話ターン数」を超えた場合にのみこの方式を適用します。リクエスト送信前にも、コンテキストのトークン数がモデルのウィンドウ上限に近づいた場合、同じ方式でコンテキストを縮小し、リクエストがモデルの上限を超えないようにします。" }, + "enable_manual_context_compression": { + "description": "手動コンテキスト圧縮(実験的)", + "hint": "有効にすると、/compact は LLM を使って現在のコンテキストを要約します。要約では詳細、ロールの状態、物語上の事実が抜け落ちる可能性があります。圧縮に失敗した場合は元の履歴が保持されます。" + }, "instruction": { "description": "コンテキスト圧縮用プロンプト", "hint": "空欄の場合はデフォルトのプロンプトを使用します。" diff --git a/tests/unit/test_agent_runner_config.py b/tests/unit/test_agent_runner_config.py index 319ed18303..d0f95f16d9 100644 --- a/tests/unit/test_agent_runner_config.py +++ b/tests/unit/test_agent_runner_config.py @@ -200,9 +200,14 @@ def test_local_legacy_fields_are_fully_migrated(): }.intersection(config["provider_settings"]) -def test_local_migration_replaces_default_root_inserted_before_version_bump(): +@pytest.mark.parametrize("default_max_steps", [30, 128]) +@pytest.mark.parametrize("manual_compression", [None, False, True]) +def test_local_migration_replaces_default_root_inserted_before_version_bump( + default_max_steps, manual_compression +): legacy_default_config = get_agent_runner_config_default("local") legacy_default_config["compression"].pop("enable_manual_context_compression") + legacy_default_config["misc"]["max_steps"] = default_max_steps config = { "config_version": 2, "provider": [ @@ -230,6 +235,10 @@ def test_local_migration_replaces_default_root_inserted_before_version_bump(): "config": legacy_default_config, }, } + if manual_compression is not None: + config["provider_settings"]["enable_manual_context_compression"] = ( + manual_compression + ) default_config = { "provider": [ @@ -239,6 +248,7 @@ def test_local_migration_replaces_default_root_inserted_before_version_bump(): ] } assert _migrate_agent_runner_config(config, default_config) + assert config["config_version"] == 4 assert config["agent_runner"] == { "runner_type": "local", @@ -257,7 +267,7 @@ def test_local_migration_replaces_default_root_inserted_before_version_bump(): "max_turns": 24, "trim_turns": 4, "overflow_strategy": "llm_compress", - "enable_manual_context_compression": False, + "enable_manual_context_compression": manual_compression is True, "instruction": "Keep decisions", "keep_recent_ratio": 0.15, "provider_id": "compressor", @@ -280,8 +290,10 @@ def test_local_migration_replaces_default_root_inserted_before_version_bump(): "llm_compress_provider_id", "tool_call_timeout", "sanitize_context_by_modalities", + "enable_manual_context_compression", } ) + assert not _migrate_agent_runner_config(config, default_config) def test_missing_local_provider_references_are_removed(): @@ -571,8 +583,12 @@ def test_agent_step_limit_upgrade_is_persisted_once( if config_version == 2: config.pop("agent_runner") config["provider_settings"]["max_agent_step"] = max_steps + config["provider_settings"]["enable_manual_context_compression"] = True else: config["agent_runner"]["config"]["misc"]["max_steps"] = max_steps + config["agent_runner"]["config"]["compression"][ + "enable_manual_context_compression" + ] = True config_path.write_text(json.dumps(config), encoding="utf-8") loaded = AstrBotConfig(config_path=str(config_path)) @@ -581,12 +597,24 @@ def test_agent_step_limit_upgrade_is_persisted_once( persisted = json.loads(config_path.read_text(encoding="utf-8-sig")) assert persisted["config_version"] == max(config_version, 4) assert persisted["agent_runner"]["config"]["misc"]["max_steps"] == expected_steps + assert ( + persisted["agent_runner"]["config"]["compression"][ + "enable_manual_context_compression" + ] + is True + ) loaded["agent_runner"]["config"]["misc"]["max_steps"] = 30 loaded.save_config() for _ in range(2): reloaded = AstrBotConfig(config_path=str(config_path)) assert reloaded["agent_runner"]["config"]["misc"]["max_steps"] == 30 + assert ( + reloaded["agent_runner"]["config"]["compression"][ + "enable_manual_context_compression" + ] + is True + ) assert not _migrate_agent_runner_config(reloaded) diff --git a/tests/unit/test_config.py b/tests/unit/test_config.py index 061ff51f1f..b6d7b81768 100644 --- a/tests/unit/test_config.py +++ b/tests/unit/test_config.py @@ -998,7 +998,7 @@ def test_nested_object_schema(self, temp_config_path): class TestConfigMetadataI18n: """Tests for i18n utils.""" - @pytest.mark.parametrize("locale", ["en-US", "zh-CN", "ru-RU"]) + @pytest.mark.parametrize("locale", ["en-US", "zh-CN", "ru-RU", "ja-JP"]) def test_manual_compression_metadata_uses_translated_runner_config_keys( self, locale: str, From f6dc857b2830c7a32bd28074d1ee6fdba79e419a Mon Sep 17 00:00:00 2001 From: C10H14N2O5 <100066858+C10H14N2O5@users.noreply.github.com> Date: Sun, 20 Sep 2026 01:03:39 +0800 Subject: [PATCH 15/16] fix: preserve token counter compatibility --- astrbot/core/agent/context/manager.py | 3 +- astrbot/core/agent/context/token_counter.py | 33 ++++++++++++++++-- tests/agent/test_context_manager.py | 18 ++++++++-- tests/agent/test_token_counter.py | 38 ++++++++++++++++++--- 4 files changed, 81 insertions(+), 11 deletions(-) diff --git a/astrbot/core/agent/context/manager.py b/astrbot/core/agent/context/manager.py index cb96fdf88c..c9c548fbf6 100644 --- a/astrbot/core/agent/context/manager.py +++ b/astrbot/core/agent/context/manager.py @@ -82,8 +82,7 @@ async def process( # 2. 基于 token 的压缩 if self.config.max_context_tokens > 0: threshold_tokens = self.token_counter.count_tokens( - result, - reported_token_usage=trusted_token_usage, + result, trusted_token_usage ) if self.compressor.should_compress( diff --git a/astrbot/core/agent/context/token_counter.py b/astrbot/core/agent/context/token_counter.py index 696e219959..8e5f6bb358 100644 --- a/astrbot/core/agent/context/token_counter.py +++ b/astrbot/core/agent/context/token_counter.py @@ -13,13 +13,13 @@ class TokenCounter(Protocol): """ def count_tokens( - self, messages: list[Message], reported_token_usage: int = 0 + self, messages: list[Message], trusted_token_usage: int = 0 ) -> int: """Count the total tokens in the message list. Args: messages: The message list. - reported_token_usage: The total token usage that LLM API returned. + trusted_token_usage: The total token usage that LLM API returned. For some cases, this value is more accurate. But some API does not return it, so the value defaults to 0. @@ -54,8 +54,35 @@ class EstimateTokenCounter: """ def count_tokens( - self, messages: list[Message], reported_token_usage: int = 0 + self, + messages: list[Message], + reported_token_usage: int = 0, + **legacy_usage: int, ) -> int: + """Use positive provider usage or estimate tokens from the messages. + + Args: + messages: The message list to estimate when usage is unavailable. + reported_token_usage: Provider usage, preferred when positive. + **legacy_usage: Accepts only the legacy ``trusted_token_usage`` keyword, + used when ``reported_token_usage`` is nonpositive. + + Returns: + Positive provider usage, or the estimated message token count. + + Raises: + TypeError: An unsupported keyword argument was supplied. + """ + unexpected = legacy_usage.keys() - {"trusted_token_usage"} + if unexpected: + name = next(iter(unexpected)) + raise TypeError( + f"count_tokens() got an unexpected keyword argument '{name}'" + ) + if reported_token_usage <= 0: + reported_token_usage = legacy_usage.get( + "trusted_token_usage", reported_token_usage + ) if reported_token_usage > 0: return reported_token_usage diff --git a/tests/agent/test_context_manager.py b/tests/agent/test_context_manager.py index aa46158acf..e25b03a444 100644 --- a/tests/agent/test_context_manager.py +++ b/tests/agent/test_context_manager.py @@ -12,6 +12,7 @@ from astrbot.core.agent.context.config import ContextConfig from astrbot.core.agent.context.manager import ContextManager +from astrbot.core.agent.context.token_counter import EstimateTokenCounter from astrbot.core.agent.message import AudioURLPart, ImageURLPart, Message, TextPart from astrbot.core.provider.entities import LLMResponse @@ -635,11 +636,22 @@ def mock_should_compress(*args, **kwargs): assert len(result) <= len(messages) @pytest.mark.asyncio + @pytest.mark.parametrize("use_legacy_counter", [False, True]) async def test_trusted_usage_triggers_compression_before_provider_call( - self, caplog + self, caplog, use_legacy_counter ): + class LegacyTokenCounter: + def count_tokens(self, messages, trusted_token_usage=0): + if trusted_token_usage > 0: + return trusted_token_usage + return EstimateTokenCounter().count_tokens(messages) + caplog.set_level("INFO", logger="astrbot") - config = ContextConfig(max_context_tokens=100, truncate_turns=1) + config = ContextConfig( + max_context_tokens=100, + truncate_turns=1, + custom_token_counter=LegacyTokenCounter() if use_legacy_counter else None, + ) manager = ContextManager(config) messages = [self.create_message("user", "short")] compressed = [self.create_message("user", "compressed")] @@ -655,6 +667,8 @@ async def test_trusted_usage_triggers_compression_before_provider_call( assert result == compressed assert "Compress completed." in caplog.text assert " 83 ->" not in caplog.text + assert " 1 -> 3 tokens" in caplog.text + assert "Context processing failed" not in caplog.text @pytest.mark.asyncio async def test_force_compression_bypasses_automatic_guards(self): diff --git a/tests/agent/test_token_counter.py b/tests/agent/test_token_counter.py index 6375b45452..61fc413959 100644 --- a/tests/agent/test_token_counter.py +++ b/tests/agent/test_token_counter.py @@ -1,5 +1,7 @@ """Tests for EstimateTokenCounter multimodal support.""" +import pytest + from astrbot.core.agent.context.token_counter import ( AUDIO_TOKEN_ESTIMATE, IMAGE_TOKEN_ESTIMATE, @@ -128,8 +130,27 @@ def test_plain_text_estimate_unchanged(self): class TestReportedUsage: - def test_reported_overrides(self): - """如果 API 返回了 token 数,直接用它不做估算。""" + @pytest.mark.parametrize( + ("args", "usage", "expected"), + [ + ((42,), {}, 42), + ((), {"reported_token_usage": 42}, 42), + ((), {"trusted_token_usage": 42}, 42), + ((), {"reported_token_usage": 42, "trusted_token_usage": 99}, 42), + ((), {"reported_token_usage": 0, "trusted_token_usage": 42}, 42), + ((), {"reported_token_usage": -1, "trusted_token_usage": 42}, 42), + ((), {}, None), + ((0,), {}, None), + ((-1,), {}, None), + ((), {"reported_token_usage": 0}, None), + ((), {"reported_token_usage": -1}, None), + ((), {"trusted_token_usage": 0}, None), + ((), {"trusted_token_usage": -1}, None), + ((), {"reported_token_usage": -1, "trusted_token_usage": 0}, None), + ], + ) + def test_reported_overrides(self, args, usage, expected): + """Positive usage overrides estimates through both supported names.""" msg = _msg( "user", [ @@ -139,8 +160,17 @@ def test_reported_overrides(self): ), ], ) - tokens = counter.count_tokens([msg], reported_token_usage=42) - assert tokens == 42 + tokens = counter.count_tokens([msg], *args, **usage) + assert tokens == (counter.count_tokens([msg]) if expected is None else expected) + + @pytest.mark.parametrize("reported_token_usage", [0, 42]) + def test_unknown_keyword_rejected(self, reported_token_usage): + with pytest.raises( + TypeError, match="unexpected keyword argument 'token_usage'" + ): + counter.count_tokens( + [], reported_token_usage=reported_token_usage, token_usage=42 + ) class TestToolCalls: From 96523220ad509cf4b96365a3211be94d38c2f7af Mon Sep 17 00:00:00 2001 From: C10H14N2O5 <100066858+C10H14N2O5@users.noreply.github.com> Date: Sun, 20 Sep 2026 01:13:35 +0800 Subject: [PATCH 16/16] fix: retain upstream token counter interface --- astrbot/core/agent/context/token_counter.py | 33 ++------------------- tests/agent/test_token_counter.py | 22 ++------------ 2 files changed, 6 insertions(+), 49 deletions(-) diff --git a/astrbot/core/agent/context/token_counter.py b/astrbot/core/agent/context/token_counter.py index 8e5f6bb358..31c274af64 100644 --- a/astrbot/core/agent/context/token_counter.py +++ b/astrbot/core/agent/context/token_counter.py @@ -54,37 +54,10 @@ class EstimateTokenCounter: """ def count_tokens( - self, - messages: list[Message], - reported_token_usage: int = 0, - **legacy_usage: int, + self, messages: list[Message], trusted_token_usage: int = 0 ) -> int: - """Use positive provider usage or estimate tokens from the messages. - - Args: - messages: The message list to estimate when usage is unavailable. - reported_token_usage: Provider usage, preferred when positive. - **legacy_usage: Accepts only the legacy ``trusted_token_usage`` keyword, - used when ``reported_token_usage`` is nonpositive. - - Returns: - Positive provider usage, or the estimated message token count. - - Raises: - TypeError: An unsupported keyword argument was supplied. - """ - unexpected = legacy_usage.keys() - {"trusted_token_usage"} - if unexpected: - name = next(iter(unexpected)) - raise TypeError( - f"count_tokens() got an unexpected keyword argument '{name}'" - ) - if reported_token_usage <= 0: - reported_token_usage = legacy_usage.get( - "trusted_token_usage", reported_token_usage - ) - if reported_token_usage > 0: - return reported_token_usage + if trusted_token_usage > 0: + return trusted_token_usage total = 0 for msg in messages: diff --git a/tests/agent/test_token_counter.py b/tests/agent/test_token_counter.py index 61fc413959..600664e411 100644 --- a/tests/agent/test_token_counter.py +++ b/tests/agent/test_token_counter.py @@ -129,28 +129,21 @@ def test_plain_text_estimate_unchanged(self): assert counter.count_tokens([_msg("user", "你" * 100)]) == 60 -class TestReportedUsage: +class TestTrustedUsage: @pytest.mark.parametrize( ("args", "usage", "expected"), [ ((42,), {}, 42), - ((), {"reported_token_usage": 42}, 42), ((), {"trusted_token_usage": 42}, 42), - ((), {"reported_token_usage": 42, "trusted_token_usage": 99}, 42), - ((), {"reported_token_usage": 0, "trusted_token_usage": 42}, 42), - ((), {"reported_token_usage": -1, "trusted_token_usage": 42}, 42), ((), {}, None), ((0,), {}, None), ((-1,), {}, None), - ((), {"reported_token_usage": 0}, None), - ((), {"reported_token_usage": -1}, None), ((), {"trusted_token_usage": 0}, None), ((), {"trusted_token_usage": -1}, None), - ((), {"reported_token_usage": -1, "trusted_token_usage": 0}, None), ], ) - def test_reported_overrides(self, args, usage, expected): - """Positive usage overrides estimates through both supported names.""" + def test_usage_overrides(self, args, usage, expected): + """Positive usage overrides estimates for keyword and positional calls.""" msg = _msg( "user", [ @@ -163,15 +156,6 @@ def test_reported_overrides(self, args, usage, expected): tokens = counter.count_tokens([msg], *args, **usage) assert tokens == (counter.count_tokens([msg]) if expected is None else expected) - @pytest.mark.parametrize("reported_token_usage", [0, 42]) - def test_unknown_keyword_rejected(self, reported_token_usage): - with pytest.raises( - TypeError, match="unexpected keyword argument 'token_usage'" - ): - counter.count_tokens( - [], reported_token_usage=reported_token_usage, token_usage=42 - ) - class TestToolCalls: def test_tool_calls_counted(self):