From 8cb2af87bd12cfa9318417f43f85aa9d5920319c 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 1/5] 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 +- astrbot/core/config/default.py | 13 + .../en-US/features/config-metadata.json | 4 + .../ru-RU/features/config-metadata.json | 4 + .../zh-CN/features/config-metadata.json | 4 + tests/agent/test_context_manager.py | 144 ++++- tests/test_conversation_commands.py | 538 ++++++++++++++++++ tests/unit/test_astr_main_agent.py | 83 +++ 13 files changed, 1122 insertions(+), 31 deletions(-) diff --git a/astrbot/builtin_stars/builtin_commands/commands/conversation.py b/astrbot/builtin_stars/builtin_commands/commands/conversation.py index 3cd4ad9e06..4ec8749afc 100644 --- a/astrbot/builtin_stars/builtin_commands/commands/conversation.py +++ b/astrbot/builtin_stars/builtin_commands/commands/conversation.py @@ -1,17 +1,30 @@ +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_AGENT_RUNNER_PROVIDER_ID_KEY, 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 from .utils.rst_scene import RstScene @@ -219,6 +232,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: """创建新对话""" cfg = self.context.get_config(umo=message.unified_msg_origin) diff --git a/astrbot/builtin_stars/builtin_commands/main.py b/astrbot/builtin_stars/builtin_commands/main.py index 4c5ce3f8ca..ff9bec7ee9 100644 --- a/astrbot/builtin_stars/builtin_commands/main.py +++ b/astrbot/builtin_stars/builtin_commands/main.py @@ -61,6 +61,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 e168bbf400..a862083a55 100644 --- a/astrbot/core/astr_main_agent.py +++ b/astrbot/core/astr_main_agent.py @@ -1304,30 +1304,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: @@ -1723,8 +1725,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/astrbot/core/config/default.py b/astrbot/core/config/default.py index 942bcda65f..6a426b929b 100644 --- a/astrbot/core/config/default.py +++ b/astrbot/core/config/default.py @@ -125,6 +125,7 @@ "persona_pool": ["*"], "prompt_prefix": "{{prompt}}", "context_limit_reached_strategy": "llm_compress", # or truncate_by_turns + "enable_manual_context_compression": False, "llm_compress_instruction": ( "Based on our full conversation history, produce a concise summary of key takeaways and/or project progress.\n" "The primary goal of this summary is to enable seamless continuation of the work that follows.\n" @@ -2940,6 +2941,9 @@ "dequeue_context_length": { "type": "int", }, + "enable_manual_context_compression": { + "type": "bool", + }, "streaming_response": { "type": "bool", }, @@ -3712,6 +3716,15 @@ }, "hint": "普通会话历史仅在超过“压缩前最多保留对话轮数”后执行该策略;请求发送前也会在上下文 token 接近模型窗口时使用同一策略保护本次请求。", }, + "provider_settings.enable_manual_context_compression": { + "description": "手动上下文压缩(实验性)", + "type": "bool", + "hint": "启用后,/compact 将使用 LLM 摘要当前上下文。摘要可能遗漏细节、角色状态或叙事事实;压缩失败时将保留原历史。", + "condition": { + "provider_settings.context_limit_reached_strategy": "llm_compress", + "provider_settings.agent_runner_type": "local", + }, + }, "provider_settings.llm_compress_instruction": { "description": "上下文压缩提示词", "type": "text", 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 0e6905cc5c..3751b36907 100644 --- a/dashboard/src/i18n/locales/en-US/features/config-metadata.json +++ b/dashboard/src/i18n/locales/en-US/features/config-metadata.json @@ -269,6 +269,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." + }, "llm_compress_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 dc0bbca5d6..2982d1fed1 100644 --- a/dashboard/src/i18n/locales/ru-RU/features/config-metadata.json +++ b/dashboard/src/i18n/locales/ru-RU/features/config-metadata.json @@ -269,6 +269,10 @@ ], "hint": "Постоянная история диалога использует эту стратегию только после превышения лимита раундов. Перед каждым запросом та же стратегия может защищать текущий контекст, когда токены приближаются к окну модели." }, + "enable_manual_context_compression": { + "description": "Ручное сжатие контекста (экспериментальная функция)", + "hint": "После включения команда /compact использует LLM для создания краткого содержания текущего контекста. Сводка может упустить детали, состояние роли или факты повествования. При сбое исходная история сохраняется." + }, "llm_compress_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 089e9ba91f..6d4b0e6014 100644 --- a/dashboard/src/i18n/locales/zh-CN/features/config-metadata.json +++ b/dashboard/src/i18n/locales/zh-CN/features/config-metadata.json @@ -271,6 +271,10 @@ ], "hint": "普通会话历史仅在超过\"压缩前最多保留对话轮数\"后执行该策略;请求发送前也会在上下文 token 接近模型窗口时使用同一策略保护本次请求。" }, + "enable_manual_context_compression": { + "description": "手动上下文压缩(实验性)", + "hint": "启用后,/compact 将使用 LLM 摘要当前上下文。摘要可能遗漏细节、角色状态或叙事事实;压缩失败时将保留原历史。" + }, "llm_compress_instruction": { "description": "上下文压缩提示词", "hint": "如果为空则使用默认提示词。" 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 3e56b6acf3..d76ab99261 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, @@ -191,3 +321,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 3892a17c4b..794bd9e70c 100644 --- a/tests/unit/test_astr_main_agent.py +++ b/tests/unit/test_astr_main_agent.py @@ -350,6 +350,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 1ffd0a029249514cfab7fb526ee2f7194b7d5b11 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 2/5] 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 4ec8749afc..e314594c63 100644 --- a/astrbot/builtin_stars/builtin_commands/commands/conversation.py +++ b/astrbot/builtin_stars/builtin_commands/commands/conversation.py @@ -298,7 +298,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 0b72b582d7..95e5ce93bb 100644 --- a/astrbot/dashboard/services/chat_service.py +++ b/astrbot/dashboard/services/chat_service.py @@ -36,6 +36,7 @@ SSE_HEARTBEAT = ": heartbeat\n\n" CHAT_RUN_SUBSCRIBER_QUEUE_SIZE = 256 +WEBCHAT_EPHEMERAL_CHAIN_TYPE = "webchat_ephemeral" WEBCHAT_IMAGE_MIME_TYPES = { ".jpg": "image/jpeg", ".jpeg": "image/jpeg", @@ -971,8 +972,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, @@ -1020,7 +1026,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 6c1ea216da..71062eeed9 100644 --- a/tests/test_chat_route.py +++ b/tests/test_chat_route.py @@ -127,17 +127,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, }, ) @@ -152,8 +171,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 d76ab99261..54658871bf 100644 --- a/tests/test_conversation_commands.py +++ b/tests/test_conversation_commands.py @@ -376,7 +376,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, @@ -440,7 +440,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 @@ -605,7 +605,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 @@ -667,9 +667,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 @@ -727,5 +730,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 f9dd9c32c36179cd82148197dc7bff1b49033913 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 3/5] 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 d546dd40a9272105bdb7baa5db4100e936b9cee2 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 4/5] 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 7c60cb23ec..8b3ace321f 100644 --- a/astrbot/core/agent/context/token_counter.py +++ b/astrbot/core/agent/context/token_counter.py @@ -12,13 +12,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. @@ -44,10 +44,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 8c91adbbfd..7dc1c3a5c0 100644 --- a/astrbot/core/agent/runners/tool_loop_agent_runner.py +++ b/astrbot/core/agent/runners/tool_loop_agent_runner.py @@ -807,7 +807,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 49d44a93d4..d587e0ad14 100644 --- a/tests/agent/test_token_counter.py +++ b/tests/agent/test_token_counter.py @@ -102,8 +102,8 @@ def test_multiple_images(self): assert tokens == IMAGE_TOKEN_ESTIMATE * 3 -class TestTrustedUsage: - def test_trusted_overrides(self): +class TestReportedUsage: + def test_reported_overrides(self): """如果 API 返回了 token 数,直接用它不做估算。""" msg = _msg( "user", @@ -114,7 +114,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 1e679de4aa..fd5accbab9 100644 --- a/tests/test_tool_loop_agent_runner.py +++ b/tests/test_tool_loop_agent_runner.py @@ -565,7 +565,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 @@ -595,7 +595,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 78915fc6ffab84ed326379b8809a5466b837d55b 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 5/5] 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 71062eeed9..2b4c2534b2 100644 --- a/tests/test_chat_route.py +++ b/tests/test_chat_route.py @@ -252,6 +252,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()