From cb8dd5518f642a7d86c1012f6e5073d4a6eea7f8 Mon Sep 17 00:00:00 2001 From: mindfn Date: Tue, 10 Mar 2026 11:48:03 +0800 Subject: [PATCH 1/5] Fix baseline context compression and transport loop regressions Restore the sync context-compression API expected by the current test suite and model assembly path. This reintroduces dare_framework.compression.core, restores Context.compress() semantics, preserves retrieval-fusion metadata assembly, and keeps the moving-compressor path available through assemble_for_model(). Also fix the concurrent transport-loop contract test to use canonical Message inputs so the background task reaches the task-local event barrier instead of failing early. --- dare_framework/agent/react_agent.py | 103 ++++- dare_framework/compression/__init__.py | 9 +- dare_framework/compression/core.py | 435 ++++++++++++++++++ dare_framework/context/context.py | 240 +++++++++- dare_framework/context/kernel.py | 4 +- .../test_base_agent_transport_contract.py | 4 +- tests/unit/test_context_implementation.py | 23 + 7 files changed, 790 insertions(+), 28 deletions(-) create mode 100644 dare_framework/compression/core.py diff --git a/dare_framework/agent/react_agent.py b/dare_framework/agent/react_agent.py index b6367958..d0f2c531 100644 --- a/dare_framework/agent/react_agent.py +++ b/dare_framework/agent/react_agent.py @@ -112,6 +112,11 @@ def __init__( tool_gateway: IToolGateway, plan_provider: IToolProvider | None = None, max_tool_rounds: int = 10, + auto_compress: bool = False, + compress_trigger_ratio: float = 0.9, + compress_target_ratio: float = 0.75, + compress_max_messages: int | None = None, + compress_strategy: str = "dedup_then_truncate", agent_channel: AgentChannel | None = None, ) -> None: super().__init__(name, agent_channel=agent_channel) @@ -120,6 +125,19 @@ def __init__( self._context = context self._tool_gateway = tool_gateway self._plan_provider = plan_provider + self._auto_compress = bool(auto_compress) + self._compress_trigger_ratio = _clamp_ratio(compress_trigger_ratio, default=0.9) + self._compress_target_ratio = _clamp_ratio(compress_target_ratio, default=0.75) + self._compress_max_messages = ( + compress_max_messages + if isinstance(compress_max_messages, int) and compress_max_messages > 0 + else None + ) + self._compress_strategy = ( + compress_strategy.strip() + if isinstance(compress_strategy, str) and compress_strategy.strip() + else "dedup_then_truncate" + ) self._context.set_tool_gateway(self._tool_gateway) # 运行时检测是否可以启用 SmartContext 能力 @@ -136,7 +154,7 @@ def plan_provider(self) -> IToolProvider | None: async def execute( self, - task: Message, + task: Message | str, *, transport: AgentChannel | None = None, ) -> RunResult: @@ -147,12 +165,13 @@ async def execute( async def _execute_basic( self, - task: Message, + task: Message | str, *, transport: AgentChannel | None = None, ) -> RunResult: """原始基础 ReAct 循环实现。""" - self._context.stm_add(task) + user_message = task if isinstance(task, Message) else Message(role="user", text=task) + self._context.stm_add(user_message) gateway = self._tool_gateway @@ -164,6 +183,9 @@ async def _execute_basic( print(f"[{self.name}] Round {round_idx + 1}/{self._max_tool_rounds}: 调用模型中...", flush=True) assembled = await self._context.assemble_for_model() messages = self._build_model_messages(assembled) + if self._maybe_auto_compress(messages): + assembled = await self._context.assemble_for_model() + messages = self._build_model_messages(assembled) model_input = ModelInput( messages=messages, @@ -329,7 +351,7 @@ async def _execute_basic( async def _execute_with_smart_context( self, - task: Message, + task: Message | str, *, transport: AgentChannel | None = None, ) -> RunResult: @@ -344,7 +366,7 @@ async def _execute_with_smart_context( return await self._execute_basic(task, transport=transport) _ = transport - source_user_message = task + source_user_message = task if isinstance(task, Message) else Message(role="user", text=task) user_message = Message( role=source_user_message.role, kind=source_user_message.kind, @@ -405,6 +427,24 @@ async def _execute_with_smart_context( messages.append(self._next_round_reflection_prompt) self._next_round_reflection_prompt = None + if self._maybe_auto_compress(messages): + assembled = await self._context.assemble_for_model() + messages = list(assembled.messages) + prompt_def = getattr(assembled, "sys_prompt", None) + sys_prompt_message = ( + Message( + role=prompt_def.role, + text=prompt_def.content, + name=prompt_def.name, + metadata=dict(prompt_def.metadata), + mark=MessageMark.IMMUTABLE, + id="sys_prompt", + ) + if prompt_def is not None + else None + ) + messages = self._context.order_messages_for_llm(messages, sys_prompt_message) + # Inject critical_block from plan_provider (maintained by plan tools) # Disabled: skip injection to observe plan agent behavior without it if False and self._plan_provider is not None: @@ -608,6 +648,38 @@ def _build_model_messages(self, assembled: Any) -> list[Message]: ) return messages + def _maybe_auto_compress(self, model_messages: list[Message]) -> bool: + """Auto-compress context before model invocation when token estimate is near budget.""" + if not self._auto_compress: + return False + max_tokens = self._context.budget.max_tokens + if max_tokens is None or max_tokens <= 0: + return False + + estimated_tokens = _estimate_messages_tokens(model_messages) + trigger_tokens = max(1, int(max_tokens * self._compress_trigger_ratio)) + if estimated_tokens < trigger_tokens: + return False + + stm_messages = self._context.stm_get() + if not stm_messages: + return False + + max_messages = self._compress_max_messages + if max_messages is None: + max_messages = max(1, int(len(stm_messages) * self._compress_target_ratio)) + if max_messages >= len(stm_messages): + max_messages = max(1, len(stm_messages) - 1) + + target_tokens = max(1, int(max_tokens * self._compress_target_ratio)) + self._context.compress( + strategy=self._compress_strategy, + max_messages=max_messages, + target_tokens=target_tokens, + tool_pair_safe=True, + ) + return True + async def _emit_terminal_transport_message( self, *, @@ -749,4 +821,25 @@ def _tool_calls_signature(tool_calls: list[dict[str, Any]]) -> tuple[str, ...]: return tuple(signature) +def _estimate_messages_tokens(messages: list[Message]) -> int: + total = 0 + for message in messages: + content = (message.text or "").strip() + attachment_tokens = len(message.attachments) * 32 + total += max(1, len(content) // 4) + attachment_tokens + 8 + return total + + +def _clamp_ratio(value: Any, *, default: float) -> float: + try: + ratio = float(value) + except (TypeError, ValueError): + return default + if not math.isfinite(ratio) or ratio <= 0: + return default + if ratio > 1: + return 1.0 + return ratio + + __all__ = ["ReactAgent"] diff --git a/dare_framework/compression/__init__.py b/dare_framework/compression/__init__.py index e0b77334..c8020700 100644 --- a/dare_framework/compression/__init__.py +++ b/dare_framework/compression/__init__.py @@ -1,11 +1,8 @@ -"""Compression utilities for context and memories. - -- MovingCompressor: 移动窗口式 STM 压缩(LLM 摘要),见 moving_compression。 -""" +"""Compression utilities for context and memories.""" from __future__ import annotations +from .core import compress_context, compress_context_llm_summary from .moving_compression import MovingCompressor -__all__ = ["MovingCompressor"] - +__all__ = ["compress_context", "compress_context_llm_summary", "MovingCompressor"] diff --git a/dare_framework/compression/core.py b/dare_framework/compression/core.py new file mode 100644 index 00000000..744e65ee --- /dev/null +++ b/dare_framework/compression/core.py @@ -0,0 +1,435 @@ +"""Core context compression helpers. + +This module preserves the synchronous compression entrypoints described by the +design docs while the moving-window compressor remains available for +`assemble_for_model()` flows. +""" + +from __future__ import annotations + +from typing import TYPE_CHECKING, Any, List, Tuple + +from dare_framework.context.types import Message as CtxMessage, MessageKind, MessageMark +from dare_framework.model import ModelInput + +if TYPE_CHECKING: + from dare_framework.context.kernel import IContext + from dare_framework.context.types import Message + from dare_framework.model import IModelAdapter + + +_UNCHANGED = object() + + +def _copy_message( + message: Message, + *, + data: dict[str, Any] | None | object = _UNCHANGED, + metadata: dict[str, Any] | object = _UNCHANGED, +) -> Message: + return CtxMessage( + role=message.role, + kind=message.kind, + text=message.text, + attachments=list(message.attachments), + data=message.data if data is _UNCHANGED else data, + name=message.name, + metadata=message.metadata if metadata is _UNCHANGED else dict(metadata), + mark=getattr(message, "mark", MessageMark.TEMPORARY), + id=getattr(message, "id", None), + ) + + +def _dedup_messages(messages: List[Message]) -> Tuple[List[Message], int]: + """Lightweight de-duplication on (role, text).""" + seen: set[int] = set() + result: List[Message] = [] + removed = 0 + + for msg in messages: + key = (msg.role, msg.text) + digest = hash(key) + if digest in seen: + removed += 1 + continue + seen.add(digest) + result.append(msg) + + return result, removed + + +def _build_summary_preview( + messages: List[Message], + max_messages: int, + tail_max: int = 10, +) -> Tuple[List[Message], int]: + """Heuristic, model-free summary strategy.""" + total = len(messages) + if total <= max_messages or max_messages <= 1: + return messages, 0 + + tail_capacity = max_messages - 1 + keep_tail = min(tail_max, tail_capacity, total) + head = messages[:-keep_tail] if keep_tail > 0 else messages + tail = messages[-keep_tail:] if keep_tail > 0 else [] + + if not head: + return messages, 0 + + preview_lines: List[str] = [] + for msg in head: + content = (msg.text or "").strip() + if not content: + continue + snippet = content.replace("\n", " ") + if len(snippet) > 120: + snippet = snippet[:120] + "..." + preview_lines.append(f"{msg.role}: {snippet}") + + if not preview_lines: + return messages, 0 + + summary_text = "Conversation summary (heuristic, no LLM):\n" + "\n".join(preview_lines) + summary_message = CtxMessage( + role="system", + kind=MessageKind.SUMMARY, + text=summary_text, + metadata={"compressed": True, "strategy": "summary_preview"}, + ) + + new_messages: List[Message] = [summary_message, *tail] + removed = total - len(new_messages) + return new_messages, removed + + +def _estimate_tokens(messages: List[Message]) -> int: + """Rough token estimate using a cheap character + attachment heuristic.""" + total = 0 + for msg in messages: + content = (msg.text or "").strip() + attachment_tokens = len(msg.attachments) * 32 + total += max(1, len(content) // 4) + attachment_tokens + 8 + return total + + +def _trim_to_target_tokens(messages: List[Message], target_tokens: int | None) -> Tuple[List[Message], int]: + """Trim oldest messages until estimated token size fits target_tokens.""" + if target_tokens is None or target_tokens <= 0: + return messages, 0 + if _estimate_tokens(messages) <= target_tokens: + return messages, 0 + + trimmed = list(messages) + removed = 0 + while len(trimmed) > 1 and _estimate_tokens(trimmed) > target_tokens: + removable_idx = next( + ( + idx + for idx, message in enumerate(trimmed) + if getattr(message, "mark", MessageMark.TEMPORARY) + not in (MessageMark.IMMUTABLE, MessageMark.PERSISTENT) + ), + None, + ) + if removable_idx is None: + break + trimmed.pop(removable_idx) + removed += 1 + return trimmed, removed + + +def _extract_tool_call_ids(message: Message) -> list[str]: + """Collect tool call ids declared on an assistant message.""" + if message.role != "assistant": + return [] + tool_calls = message.data.get("tool_calls", []) if isinstance(message.data, dict) else [] + if not isinstance(tool_calls, list): + return [] + + ids: list[str] = [] + for call in tool_calls: + if not isinstance(call, dict): + continue + tool_id = call.get("id") + if isinstance(tool_id, str) and tool_id.strip(): + ids.append(tool_id.strip()) + return ids + + +def _enforce_tool_pair_safety(messages: List[Message]) -> Tuple[List[Message], int]: + """Keep tool_call/tool_result in sync so compression never leaves orphan pairs.""" + tool_result_ids = { + message.name.strip() + for message in messages + if message.role == "tool" and isinstance(message.name, str) and message.name.strip() + } + + updated_messages: list[Message] = [] + retained_call_ids: set[str] = set() + retained_idless_tool_names: set[str] = set() + changes = 0 + + for message in messages: + if message.role != "assistant": + updated_messages.append(message) + continue + raw_calls = message.data.get("tool_calls", []) if isinstance(message.data, dict) else [] + if not isinstance(raw_calls, list): + updated_messages.append(message) + continue + + filtered_calls = [] + for call in raw_calls: + if not isinstance(call, dict): + continue + tool_id = call.get("id") + if isinstance(tool_id, str) and tool_id.strip() and tool_id.strip() in tool_result_ids: + filtered_calls.append(call) + retained_call_ids.add(tool_id.strip()) + continue + if not isinstance(tool_id, str) or not tool_id.strip(): + filtered_calls.append(call) + tool_name = call.get("name") + if isinstance(tool_name, str) and tool_name.strip(): + retained_idless_tool_names.add(tool_name.strip()) + + if len(filtered_calls) != len(raw_calls): + changes += len(raw_calls) - len(filtered_calls) + updated_data = dict(message.data) if isinstance(message.data, dict) else {} + updated_data["tool_calls"] = filtered_calls + updated_messages.append(_copy_message(message, data=updated_data)) + else: + retained_call_ids.update(_extract_tool_call_ids(message)) + updated_messages.append(message) + + final_messages: list[Message] = [] + for message in updated_messages: + if message.role == "tool": + tool_id = message.name.strip() if isinstance(message.name, str) else "" + if tool_id in retained_call_ids or (tool_id and tool_id in retained_idless_tool_names): + final_messages.append(message) + continue + if tool_id: + changes += 1 + continue + final_messages.append(message) + return final_messages, changes + + +def _annotate_strategy(messages: List[Message], strategy: str) -> List[Message]: + """Attach strategy metadata to the first message when compression rewrites context.""" + if not messages: + return messages + for message in messages: + if message.metadata.get("compressed") is True: + return messages + + head = messages[0] + metadata = dict(head.metadata) + metadata["compressed"] = True + metadata.setdefault("strategy", strategy) + messages[0] = _copy_message(head, metadata=metadata) + return messages + + +def compress_context( + context: IContext, + *, + phase: str | None = None, + max_messages: int | None = None, + **options: Any, +) -> None: + """Compress short-term memory for a given context.""" + _ = phase + + target_tokens_raw = options.get("target_tokens") + target_tokens: int | None = None + if target_tokens_raw is not None: + try: + target_tokens = int(target_tokens_raw) + except (TypeError, ValueError): + target_tokens = None + + if (max_messages is None or max_messages < 0) and (target_tokens is None or target_tokens <= 0): + return + + stm_get = getattr(context, "stm_get", None) + stm_clear = getattr(context, "stm_clear", None) + stm_add = getattr(context, "stm_add", None) + if not callable(stm_get) or not callable(stm_clear) or not callable(stm_add): + return + + messages: List[Message] = list(stm_get()) + if not messages: + return + + if max_messages is None: + max_messages = len(messages) + elif max_messages < 0: + max_messages = len(messages) + + strategy = options.get("strategy", "truncate") + tool_pair_safe = bool(options.get("tool_pair_safe", False)) + + removed_total = 0 + + if strategy == "summary_preview": + messages, removed = _build_summary_preview(messages, max_messages) + removed_total += removed + + if strategy == "dedup_then_truncate": + messages, removed = _dedup_messages(messages) + removed_total += removed + + if max_messages == 0: + removed_total += len(messages) + messages = [] + elif len(messages) > max_messages: + protected = [ + message + for message in messages + if getattr(message, "mark", MessageMark.TEMPORARY) + in (MessageMark.IMMUTABLE, MessageMark.PERSISTENT) + ] + temporary = [ + message + for message in messages + if getattr(message, "mark", MessageMark.TEMPORARY) + not in (MessageMark.IMMUTABLE, MessageMark.PERSISTENT) + ] + keep_temporary = max(max_messages - len(protected), 0) + if keep_temporary <= 0: + kept_tail: list[Message] = [] + elif keep_temporary < len(temporary): + kept_tail = temporary[-keep_temporary:] + else: + kept_tail = temporary + kept_tail_refs = {id(message) for message in kept_tail} + messages = [ + message + for message in messages + if ( + getattr(message, "mark", MessageMark.TEMPORARY) + in (MessageMark.IMMUTABLE, MessageMark.PERSISTENT) + ) + or id(message) in kept_tail_refs + ] + removed_total += len(protected) + len(temporary) - len(messages) + + messages, removed = _trim_to_target_tokens(messages, target_tokens) + removed_total += removed + + if tool_pair_safe: + messages, changes = _enforce_tool_pair_safety(messages) + removed_total += changes + + if removed_total == 0: + return + + messages = _annotate_strategy(messages, str(strategy)) + stm_clear() + for msg in messages: + stm_add(msg) + + +async def compress_context_llm_summary( + context: IContext, + *, + model: IModelAdapter, + max_messages: int, + keep_tail: int = 8, + system_prompt: str | None = None, + language: str = "zh", +) -> None: + """High-level compression using the LLM to generate a semantic summary.""" + if max_messages <= 1: + return + + stm_get = getattr(context, "stm_get", None) + stm_clear = getattr(context, "stm_clear", None) + stm_add = getattr(context, "stm_add", None) + if not callable(stm_get) or not callable(stm_clear) or not callable(stm_add): + return + + messages: List[Message] = list(stm_get()) + total = len(messages) + if total <= max_messages: + return + + tail_capacity = max_messages - 1 + keep_tail_eff = min(keep_tail, tail_capacity, total) + head = messages[:-keep_tail_eff] if keep_tail_eff > 0 else messages + tail = messages[-keep_tail_eff:] if keep_tail_eff > 0 else [] + + if not head: + return + + lines: List[str] = [] + for msg in head: + content = (msg.text or "").strip() + if not content: + continue + snippet = content.replace("\n", " ") + if len(snippet) > 512: + snippet = snippet[:512] + "..." + lines.append(f"{msg.role}: {snippet}") + + if not lines: + return + + conversation_text = "\n".join(lines) + if system_prompt is None: + if language == "zh": + system_prompt = ( + "你是一个对话摘要助手,请在不丢失关键信息的前提下," + "用简洁、结构化的方式总结下面的一段历史对话。" + "可以合并重复信息,但不要编造不存在的内容。" + ) + else: + system_prompt = ( + "You are a conversation summarization assistant. " + "Produce a concise, structured summary of the following history, " + "preserving key facts and decisions. Do not invent new information." + ) + + sys_msg = CtxMessage(role="system", kind=MessageKind.SUMMARY, text=system_prompt) + user_intro = ( + "下面是一段需要被压缩的历史对话,请输出一个摘要,用于后续继续对话使用。\n\n" + "=== 历史开始 ===\n" + f"{conversation_text}\n" + "=== 历史结束 ===" + if language == "zh" + else + "Here is the conversation history that needs to be compressed. " + "Please output a summary that can be used for continuing the dialogue.\n\n" + "=== HISTORY START ===\n" + f"{conversation_text}\n" + "=== HISTORY END ===" + ) + user_msg = CtxMessage(role="user", text=user_intro) + + model_input = ModelInput( + messages=[sys_msg, user_msg], + tools=[], + metadata={"compression": "llm_summary"}, + ) + + response = await model.generate(model_input) + summary_text = (response.content or "").strip() + if not summary_text: + return + + summary_message = CtxMessage( + role="system", + kind=MessageKind.SUMMARY, + text=summary_text, + metadata={"compressed": True, "strategy": "llm_summary"}, + ) + + stm_clear() + stm_add(summary_message) + for msg in tail: + stm_add(msg) + + +__all__ = ["compress_context", "compress_context_llm_summary"] diff --git a/dare_framework/context/context.py b/dare_framework/context/context.py index 70b37cae..6c05a139 100644 --- a/dare_framework/context/context.py +++ b/dare_framework/context/context.py @@ -201,21 +201,52 @@ def list_tools(self) -> list[CapabilityDescriptor]: def assemble(self) -> AssembledContext: return self._assemble_context.assemble(self) - async def compress(self, **options: Any) -> None: - """压缩 STM:仅委托 moving_compressor.prune;无 compressor 时无操作。""" - if self._moving_compressor is None: + def compress(self, **options: Any) -> None: + """Compress context to fit within budget.""" + from dare_framework.compression.core import compress_context + + # Preserve backend STM semantics (for example SmartSTM mark-based retention) + # before applying advanced compression strategies. + compress_impl = getattr(self._short_term_memory, "compress", None) + has_advanced_options = any( + key in options + for key in ("target_tokens", "tool_pair_safe", "strategy", "phase") + ) + raw_max_messages = options.get("max_messages") + max_messages = ( + raw_max_messages + if isinstance(raw_max_messages, int) and raw_max_messages >= 0 + else None + ) + if callable(compress_impl) and not has_advanced_options: + compress_impl(max_messages=max_messages) return - # 只把与 token 预算相关的参数传给压缩器;摘要 prompt 与语言策略由压缩器内部自行决定。 - prune_opts: dict[str, Any] = {} - if "max_context_tokens" in options: - prune_opts["max_context_tokens"] = options["max_context_tokens"] - elif self._context_window_tokens is not None and self._context_window_tokens > 0: - prune_opts["max_context_tokens"] = self._context_window_tokens - await self._moving_compressor.prune(self, **prune_opts) + + compress_context(self, **options) + + # For advanced compression, run backend max-message retention after strategy + # execution so strategy implementations can inspect full pre-trim history. + if callable(compress_impl) and has_advanced_options and max_messages is not None: + compress_impl(max_messages=max_messages) async def assemble_for_model(self, **options: Any) -> AssembledContext: - """供模型调用的装配入口:在内部静默触发压缩,然后返回 AssembledContext。""" - await self.compress(**options) + """Model-facing assembly path with optional moving-window compression.""" + sync_compress_keys = ("max_messages", "target_tokens", "tool_pair_safe", "strategy", "phase") + has_sync_compress_options = any(key in options for key in sync_compress_keys) + + # Avoid calling compress({}) on the default assembly path. Some callers + # override compress() for observability, and an empty no-op call changes + # behavior without providing any trimming value. + if has_sync_compress_options: + self.compress(**options) + + if self._moving_compressor is not None: + prune_opts: dict[str, Any] = {} + if "max_context_tokens" in options: + prune_opts["max_context_tokens"] = options["max_context_tokens"] + elif self._context_window_tokens is not None and self._context_window_tokens > 0: + prune_opts["max_context_tokens"] = self._context_window_tokens + await self._moving_compressor.prune(self, **prune_opts) return self.assemble() @@ -223,6 +254,81 @@ class DefaultAssembledContext(IAssembleContext): """Default context assembly strategy. """ + _DEFAULT_TOP_K = 3 + _DEFAULT_RESERVE_TOKENS = 256 + _DEFAULT_SOURCE_RATIO = 0.5 + + def _safe_int(self, value: Any, default: int, *, minimum: int | None = None) -> int: + try: + parsed = int(value) + except (TypeError, ValueError, OverflowError): + parsed = default + if minimum is not None: + parsed = max(minimum, parsed) + return parsed + + def _safe_ratio(self, value: Any) -> float: + try: + parsed = float(value) + except (TypeError, ValueError, OverflowError): + return self._DEFAULT_SOURCE_RATIO + if not math.isfinite(parsed) or parsed < 0: + return self._DEFAULT_SOURCE_RATIO + return parsed + + def _estimate_tokens(self, messages: list[Message]) -> int: + """Rough token estimate using character + attachment heuristics.""" + total = 0 + for message in messages: + content_tokens = max(1, len((message.text or "").strip()) // 4) + attachment_tokens = len(message.attachments) * 32 + total += content_tokens + attachment_tokens + 8 + return total + + def _take_with_budget(self, messages: list[Message], budget_tokens: float) -> list[Message]: + if budget_tokens == float("inf"): + return list(messages) + if budget_tokens <= 0: + return [] + + kept: list[Message] = [] + used = 0 + for message in messages: + message_tokens = self._estimate_tokens([message]) + if used + message_tokens > budget_tokens: + # Skip oversized candidates so smaller later hits can still fit. + continue + kept.append(message) + used += message_tokens + return kept + + def _derive_query(self, messages: list[Message]) -> str: + for message in reversed(messages): + if message.role == "user" and (message.text or "").strip(): + return (message.text or "").strip() + for message in reversed(messages): + if (message.text or "").strip(): + return (message.text or "").strip() + return "" + + def _load_source_options(self, config_map: dict[str, Any]) -> tuple[int, float]: + top_k = self._safe_int( + config_map.get("assemble_top_k"), + self._DEFAULT_TOP_K, + minimum=0, + ) + ratio = self._safe_ratio(config_map.get("assemble_ratio", self._DEFAULT_SOURCE_RATIO)) + return top_k, ratio + + def _set_degrade( + self, + retrieval_metadata: dict[str, Any], + *, + reason: str, + ) -> None: + retrieval_metadata["degraded"] = True + retrieval_metadata["degrade_reason"] = reason + def assemble(self, context: IContext) -> AssembledContext: messages = context.stm_get() tools = context.list_tools() @@ -232,12 +338,120 @@ def assemble(self, context: IContext) -> AssembledContext: sys_prompt = enrich_prompt_with_skill(sys_prompt, context.sys_skill) + query = self._derive_query(messages) + ltm_config = context.config.long_term_memory if isinstance(context.config.long_term_memory, dict) else {} + knowledge_config = context.config.knowledge if isinstance(context.config.knowledge, dict) else {} + ltm_top_k, ltm_ratio = self._load_source_options(ltm_config) + knowledge_top_k, knowledge_ratio = self._load_source_options(knowledge_config) + ltm_active = context.long_term_memory is not None and ltm_top_k > 0 + knowledge_active = context.knowledge is not None and knowledge_top_k > 0 + + if ltm_active and not knowledge_active: + reserve_tokens_raw = ltm_config.get("assemble_reserve_tokens") + elif knowledge_active and not ltm_active: + reserve_tokens_raw = knowledge_config.get("assemble_reserve_tokens") + else: + reserve_tokens_raw = ltm_config.get("assemble_reserve_tokens") + if reserve_tokens_raw is None: + reserve_tokens_raw = knowledge_config.get("assemble_reserve_tokens") + reserve_tokens = self._safe_int( + reserve_tokens_raw, + self._DEFAULT_RESERVE_TOKENS, + minimum=0, + ) + + retrieval_metadata: dict[str, Any] = { + "query": query, + "stm_count": len(messages), + "ltm_requested": ltm_top_k, + "knowledge_requested": knowledge_top_k, + "ltm_count": 0, + "knowledge_count": 0, + "degraded": False, + "degrade_reason": None, + } + + remaining_tokens = context.budget_remaining("tokens") + stm_token_estimate = self._estimate_tokens(messages) + retrieval_budget: float = float("inf") + if remaining_tokens != float("inf"): + retrieval_budget = max(0.0, float(remaining_tokens) - float(stm_token_estimate) - float(reserve_tokens)) + + ltm_messages: list[Message] = [] + knowledge_messages: list[Message] = [] + + if retrieval_budget <= 0 and (ltm_active or knowledge_active): + self._set_degrade(retrieval_metadata, reason="token_budget_low") + else: + ratio_total = 0.0 + if ltm_active: + ratio_total += ltm_ratio + if knowledge_active: + ratio_total += knowledge_ratio + if ratio_total <= 0: + # Fall back only across active retrieval sources. + ratio_total = 0.0 + if ltm_active: + ltm_ratio = self._DEFAULT_SOURCE_RATIO + ratio_total += ltm_ratio + if knowledge_active: + knowledge_ratio = self._DEFAULT_SOURCE_RATIO + ratio_total += knowledge_ratio + + normalized_ltm_ratio = (ltm_ratio / ratio_total) if ltm_active and ratio_total > 0 else 0.0 + normalized_knowledge_ratio = ( + (knowledge_ratio / ratio_total) if knowledge_active and ratio_total > 0 else 0.0 + ) + + ltm_budget = float("inf") + knowledge_budget = float("inf") + if retrieval_budget != float("inf"): + ltm_budget = retrieval_budget * normalized_ltm_ratio + knowledge_budget = retrieval_budget * normalized_knowledge_ratio + + ltm_retrieval_failed = False + if ltm_active: + if ltm_budget <= 0: + ltm_messages = [] + else: + try: + ltm_candidates = context.long_term_memory.get(query=query, top_k=ltm_top_k) + ltm_messages = self._take_with_budget(ltm_candidates, ltm_budget) + if len(ltm_messages) < len(ltm_candidates): + self._set_degrade(retrieval_metadata, reason="token_budget_low") + except Exception: + ltm_retrieval_failed = True + self._set_degrade(retrieval_metadata, reason="ltm_retrieval_failed") + ltm_messages = [] + + if knowledge_active: + try: + effective_knowledge_budget = knowledge_budget + if retrieval_budget != float("inf") and ltm_active and ltm_retrieval_failed: + effective_knowledge_budget = retrieval_budget + if effective_knowledge_budget <= 0: + knowledge_messages = [] + else: + knowledge_candidates = context.knowledge.get(query=query, top_k=knowledge_top_k) + knowledge_messages = self._take_with_budget(knowledge_candidates, effective_knowledge_budget) + if len(knowledge_messages) < len(knowledge_candidates): + self._set_degrade(retrieval_metadata, reason="token_budget_low") + except Exception: + if not retrieval_metadata["degraded"]: + self._set_degrade(retrieval_metadata, reason="knowledge_retrieval_failed") + knowledge_messages = [] + + merged_messages = [*messages, *ltm_messages, *knowledge_messages] + retrieval_metadata["ltm_count"] = len(ltm_messages) + retrieval_metadata["knowledge_count"] = len(knowledge_messages) + return AssembledContext( - messages=list(messages), + messages=merged_messages, sys_prompt=sys_prompt, tools=tools, metadata={ "context_id": context.id, + "retrieval": retrieval_metadata, }, ) diff --git a/dare_framework/context/kernel.py b/dare_framework/context/kernel.py index 50193a25..25b2cc48 100644 --- a/dare_framework/context/kernel.py +++ b/dare_framework/context/kernel.py @@ -104,9 +104,9 @@ def list_tools(self) -> list[CapabilityDescriptor]: ... def assemble(self) -> AssembledContext: ... - # Compress (core):由具体 Context 实现决定何时触发;默认在 assemble_for_model 中静默调用。 + # Compress (core):同步高级压缩入口;assemble_for_model 可在内部追加异步 moving compression。 - async def compress(self, **options: Any) -> None: ... + def compress(self, **options: Any) -> None: ... # Assemble for model: 默认直接调用 assemble,由具体实现决定是否在内部触发 compress。 diff --git a/tests/unit/test_base_agent_transport_contract.py b/tests/unit/test_base_agent_transport_contract.py index a3bceb9f..a51e4407 100644 --- a/tests/unit/test_base_agent_transport_contract.py +++ b/tests/unit/test_base_agent_transport_contract.py @@ -431,14 +431,14 @@ async def test_transport_loop_flag_is_task_local_for_concurrent_execute_calls() polled_task = asyncio.create_task( agent._execute_polled_message( - "loop-task", + Message(role="user", text="loop-task"), channel=channel, envelope_id="req_1", ) ) await agent.loop_execution_started.wait() - await agent.execute("direct-task", transport=channel) + await agent.execute(Message(role="user", text="direct-task"), transport=channel) agent.allow_loop_execution_finish.set() await polled_task diff --git a/tests/unit/test_context_implementation.py b/tests/unit/test_context_implementation.py index ef1b1129..ad74dcdc 100644 --- a/tests/unit/test_context_implementation.py +++ b/tests/unit/test_context_implementation.py @@ -149,6 +149,14 @@ def compress(self, **kwargs: object) -> int: return 0 +class _RecordingMovingCompressor: + def __init__(self) -> None: + self.calls: list[dict[str, object]] = [] + + async def prune(self, context: Context, **options: object) -> None: + self.calls.append({"context": context, "options": dict(options)}) + + def test_context_assemble_fuses_ltm_and_knowledge_with_latest_user_query(): ltm = _FakeRetrieval([Message(role="assistant", text="ltm-hit")]) knowledge = _FakeRetrieval([Message(role="assistant", text="knowledge-hit")]) @@ -466,6 +474,21 @@ def test_context_assemble_handles_overflowing_numeric_ratio_config() -> None: assert assembled.metadata["retrieval"]["ltm_count"] == 1 +@pytest.mark.asyncio +async def test_context_assemble_for_model_runs_moving_compressor_with_context_window_tokens() -> None: + ctx = Context(config=Config(), context_window_tokens=256) + compressor = _RecordingMovingCompressor() + ctx.set_moving_compressor(compressor) + ctx.stm_add(Message(role="user", text="query")) + + assembled = await ctx.assemble_for_model() + + assert len(compressor.calls) == 1 + assert compressor.calls[0]["context"] is ctx + assert compressor.calls[0]["options"] == {"max_context_tokens": 256} + assert [message.text for message in assembled.messages] == ["query"] + + def test_context_compress_max_messages_uses_backend_compress_only(monkeypatch: pytest.MonkeyPatch) -> None: stm = _CompressionRecordingSTM( [ From 56f6e572eed29d1de03216d9aeb4a8325d2280b1 Mon Sep 17 00:00:00 2001 From: mindfn Date: Tue, 10 Mar 2026 14:09:55 +0800 Subject: [PATCH 2/5] Address PR 211 review findings Preserve structured tool-call payloads during deduplication by incorporating message kind, name, attachments, and data into the compression identity key, so auto-compression no longer drops distinct assistant tool-call history with empty text. Also re-append the transient SmartContext reflection prompt after re-assembling messages for auto-compression, and add regression coverage for both review findings. --- dare_framework/agent/react_agent.py | 2 + dare_framework/compression/core.py | 37 ++++++++++--- tests/unit/test_context_compression.py | 50 +++++++++++++++++ .../test_react_agent_gateway_injection.py | 53 +++++++++++++++++++ 4 files changed, 136 insertions(+), 6 deletions(-) diff --git a/dare_framework/agent/react_agent.py b/dare_framework/agent/react_agent.py index d0f2c531..0a28f568 100644 --- a/dare_framework/agent/react_agent.py +++ b/dare_framework/agent/react_agent.py @@ -444,6 +444,8 @@ async def _execute_with_smart_context( else None ) messages = self._context.order_messages_for_llm(messages, sys_prompt_message) + if injected_reflection_prompt is not None: + messages.append(injected_reflection_prompt) # Inject critical_block from plan_provider (maintained by plan tools) # Disabled: skip injection to observe plan agent behavior without it diff --git a/dare_framework/compression/core.py b/dare_framework/compression/core.py index 744e65ee..531e4775 100644 --- a/dare_framework/compression/core.py +++ b/dare_framework/compression/core.py @@ -21,6 +21,25 @@ _UNCHANGED = object() +def _freeze_value(value: Any) -> Any: + """Build a hashable structural key for nested message payloads.""" + if isinstance(value, dict): + return tuple(sorted((str(key), _freeze_value(item)) for key, item in value.items())) + if isinstance(value, list): + return tuple(_freeze_value(item) for item in value) + if isinstance(value, tuple): + return tuple(_freeze_value(item) for item in value) + if hasattr(value, "kind") and hasattr(value, "uri"): + return ( + getattr(value.kind, "value", value.kind), + value.uri, + value.mime_type, + value.filename, + _freeze_value(getattr(value, "metadata", {})), + ) + return value + + def _copy_message( message: Message, *, @@ -41,18 +60,24 @@ def _copy_message( def _dedup_messages(messages: List[Message]) -> Tuple[List[Message], int]: - """Lightweight de-duplication on (role, text).""" - seen: set[int] = set() + """De-duplicate only when the full public message payload matches.""" + seen: set[Any] = set() result: List[Message] = [] removed = 0 for msg in messages: - key = (msg.role, msg.text) - digest = hash(key) - if digest in seen: + key = ( + msg.role, + msg.kind, + msg.text, + msg.name, + _freeze_value(msg.attachments), + _freeze_value(msg.data), + ) + if key in seen: removed += 1 continue - seen.add(digest) + seen.add(key) result.append(msg) return result, removed diff --git a/tests/unit/test_context_compression.py b/tests/unit/test_context_compression.py index 1576fdad..73c626b8 100644 --- a/tests/unit/test_context_compression.py +++ b/tests/unit/test_context_compression.py @@ -142,6 +142,56 @@ def test_compress_context_tool_pair_safe_preserves_assistant_id_and_mark_when_fi } +def test_compress_context_dedup_preserves_distinct_tool_call_payloads() -> None: + ctx = Context(config=Config()) + ctx.stm_add( + Message( + role="assistant", + kind=MessageKind.TOOL_CALL, + text="", + data={"tool_calls": [{"id": "tc_1", "name": "demo_tool", "arguments": {"x": 1}}]}, + ) + ) + ctx.stm_add( + Message( + role="tool", + kind=MessageKind.TOOL_RESULT, + name="tc_1", + text='{"success": true}', + data={"success": True}, + ) + ) + ctx.stm_add( + Message( + role="assistant", + kind=MessageKind.TOOL_CALL, + text="", + data={"tool_calls": [{"id": "tc_2", "name": "demo_tool", "arguments": {"x": 2}}]}, + ) + ) + ctx.stm_add( + Message( + role="tool", + kind=MessageKind.TOOL_RESULT, + name="tc_2", + text='{"success": true}', + data={"success": True}, + ) + ) + + compress_context(ctx, strategy="dedup_then_truncate", max_messages=10, tool_pair_safe=True) + + assistant_tool_ids = [ + _tool_ids(message) + for message in ctx.stm_get() + if message.role == "assistant" and message.kind == MessageKind.TOOL_CALL + ] + tool_result_ids = [message.name for message in ctx.stm_get() if message.role == "tool"] + + assert assistant_tool_ids == [["tc_1"], ["tc_2"]] + assert tool_result_ids == ["tc_1", "tc_2"] + + def test_compress_context_target_tokens_trims_long_history() -> None: ctx = Context(config=Config()) for idx in range(8): diff --git a/tests/unit/test_react_agent_gateway_injection.py b/tests/unit/test_react_agent_gateway_injection.py index ff90142e..07ebc2f1 100644 --- a/tests/unit/test_react_agent_gateway_injection.py +++ b/tests/unit/test_react_agent_gateway_injection.py @@ -7,6 +7,7 @@ from dare_framework.agent.react_agent import ReactAgent from dare_framework.config import Config from dare_framework.context import Context +from dare_framework.context.manage_context import MANAGE_CONTEXT_TOOL_NAME from dare_framework.context.types import MessageKind from dare_framework.context.smartcontext import SmartContext from dare_framework.model.types import ModelInput, ModelResponse @@ -178,6 +179,16 @@ async def generate(self, model_input: ModelInput, *, options: Any | None = None) return ModelResponse(content="final", tool_calls=[]) +class _CapturingFinalModel: + def __init__(self) -> None: + self.last_messages: list[Any] | None = None + + async def generate(self, model_input: ModelInput, *, options: Any | None = None) -> ModelResponse: + _ = options + self.last_messages = list(model_input.messages) + return ModelResponse(content="final", tool_calls=[]) + + class _NonConvergingToolModel: def __init__(self) -> None: self._idx = 0 @@ -197,6 +208,21 @@ async def generate(self, model_input: ModelInput, *, options: Any | None = None) ) +class _ManageContextGateway(_RecordingGateway): + def __init__(self) -> None: + super().__init__("manage-context") + self._capabilities = [ + CapabilityDescriptor( + id=MANAGE_CONTEXT_TOOL_NAME, + type=CapabilityType.TOOL, + name=MANAGE_CONTEXT_TOOL_NAME, + description="manage context", + input_schema={"type": "object"}, + output_schema={"type": "object"}, + ) + ] + + @pytest.mark.asyncio async def test_react_agent_prefers_injected_gateway_over_context_gateway() -> None: context_gateway = _RecordingGateway("context") @@ -456,6 +482,33 @@ async def test_react_agent_auto_compress_triggers_in_smart_context_path() -> Non assert context.compress_calls[0].get("tool_pair_safe") is True +@pytest.mark.asyncio +async def test_react_agent_auto_compress_reappends_reflection_prompt_in_smart_context_path() -> None: + context = _CompressionRecordingSmartContext(config=Config()) + context.budget.max_tokens = 100 + context.update_task_complete(True) + gateway = _ManageContextGateway() + model = _CapturingFinalModel() + agent = ReactAgent( + name="react-test-smartcontext-reflection-prompt", + model=model, + context=context, + tool_gateway=gateway, + auto_compress=True, + compress_trigger_ratio=0.01, + compress_target_ratio=0.5, + ) + + result = await agent("smart context compress") + + assert result.success is True + assert model.last_messages is not None + assert any( + (message.text or "") == "【提示】请先调用 manage_context 根据任务初始化 context 状态。" + for message in model.last_messages + ) + + @pytest.mark.asyncio async def test_react_agent_loop_guard_emits_terminal_message_event() -> None: context = Context(config=Config()) From a6aff312adeccd61d05c11fb045d890465a5547e Mon Sep 17 00:00:00 2001 From: mindfn Date: Tue, 10 Mar 2026 14:41:45 +0800 Subject: [PATCH 3/5] Fix zero-message backend compression semantics Make the default in-memory STM treat max_messages=0 as full eviction instead of relying on Python's -0 slice behavior, which previously retained the entire history. Add a Context-level regression test so the canonical basic compression path guarantees explicit zero-message requests clear the default STM. --- dare_framework/memory/in_memory_stm.py | 8 +++++++- tests/unit/test_context_implementation.py | 10 ++++++++++ 2 files changed, 17 insertions(+), 1 deletion(-) diff --git a/dare_framework/memory/in_memory_stm.py b/dare_framework/memory/in_memory_stm.py index a5e95c62..8035d1f4 100644 --- a/dare_framework/memory/in_memory_stm.py +++ b/dare_framework/memory/in_memory_stm.py @@ -50,7 +50,13 @@ def compress(self, max_messages: int | None = None, **kwargs) -> int: Returns: Number of messages removed. """ - if max_messages is None or len(self._messages) <= max_messages: + if max_messages is None: + return 0 + if max_messages <= 0: + removed_count = len(self._messages) + self._messages = [] + return removed_count + if len(self._messages) <= max_messages: return 0 removed_count = len(self._messages) - max_messages diff --git a/tests/unit/test_context_implementation.py b/tests/unit/test_context_implementation.py index ad74dcdc..2ae37a08 100644 --- a/tests/unit/test_context_implementation.py +++ b/tests/unit/test_context_implementation.py @@ -515,6 +515,16 @@ def _unexpected_compress_context(*args: object, **kwargs: object) -> None: assert [message.text for message in ctx.stm_get()] == ["m1", "m2"] +def test_context_compress_zero_max_messages_clears_default_stm() -> None: + ctx = Context(config=Config()) + ctx.stm_add(Message(role="user", text="m0")) + ctx.stm_add(Message(role="assistant", text="m1")) + + ctx.compress(max_messages=0) + + assert ctx.stm_get() == [] + + def test_context_compress_advanced_path_preserves_backend_semantics(monkeypatch: pytest.MonkeyPatch) -> None: stm = _CompressionRecordingSTM( [ From 73b04b28edc0c6ae28eba1b50c26ddeb0923b8eb Mon Sep 17 00:00:00 2001 From: mindfn Date: Tue, 10 Mar 2026 15:22:19 +0800 Subject: [PATCH 4/5] Harden compression and retrieval budget edge cases Make budget_remaining treat zero-valued limits as finite instead of falling back to infinity, so strict zero-token runs correctly degrade retrieval instead of fetching LTM/knowledge under an exhausted budget. Also make compression dedup keys robust for unhashable payload values, and add regressions for both the zero-budget retrieval path and non-JSON-native payload deduplication. --- dare_framework/compression/core.py | 12 ++++++++++++ dare_framework/context/context.py | 12 ++++++++---- tests/unit/test_context_compression.py | 12 ++++++++++++ tests/unit/test_context_implementation.py | 20 ++++++++++++++++++++ 4 files changed, 52 insertions(+), 4 deletions(-) diff --git a/dare_framework/compression/core.py b/dare_framework/compression/core.py index 531e4775..2273eef2 100644 --- a/dare_framework/compression/core.py +++ b/dare_framework/compression/core.py @@ -7,6 +7,7 @@ from __future__ import annotations +from dataclasses import asdict, is_dataclass from typing import TYPE_CHECKING, Any, List, Tuple from dare_framework.context.types import Message as CtxMessage, MessageKind, MessageMark @@ -29,6 +30,11 @@ def _freeze_value(value: Any) -> Any: return tuple(_freeze_value(item) for item in value) if isinstance(value, tuple): return tuple(_freeze_value(item) for item in value) + if isinstance(value, (set, frozenset)): + frozen_items = [_freeze_value(item) for item in value] + return tuple(sorted(frozen_items, key=repr)) + if is_dataclass(value) and not isinstance(value, type): + return ("dataclass", type(value).__qualname__, _freeze_value(asdict(value))) if hasattr(value, "kind") and hasattr(value, "uri"): return ( getattr(value.kind, "value", value.kind), @@ -37,6 +43,12 @@ def _freeze_value(value: Any) -> Any: value.filename, _freeze_value(getattr(value, "metadata", {})), ) + try: + hash(value) + except TypeError: + if hasattr(value, "__dict__"): + return (type(value).__qualname__, _freeze_value(vars(value))) + return (type(value).__qualname__, repr(value)) return value diff --git a/dare_framework/context/context.py b/dare_framework/context/context.py index 6c05a139..a6b5ebbf 100644 --- a/dare_framework/context/context.py +++ b/dare_framework/context/context.py @@ -176,13 +176,17 @@ def budget_remaining(self, resource: str) -> float: """Get remaining budget for a resource.""" b = self._budget if resource == "tokens": - return (b.max_tokens - b.used_tokens) if b.max_tokens else float("inf") + return (b.max_tokens - b.used_tokens) if b.max_tokens is not None else float("inf") elif resource == "cost": - return (b.max_cost - b.used_cost) if b.max_cost else float("inf") + return (b.max_cost - b.used_cost) if b.max_cost is not None else float("inf") elif resource == "tool_calls": - return (b.max_tool_calls - b.used_tool_calls) if b.max_tool_calls else float("inf") + return (b.max_tool_calls - b.used_tool_calls) if b.max_tool_calls is not None else float("inf") elif resource == "time_seconds": - return (b.max_time_seconds - b.used_time_seconds) if b.max_time_seconds else float("inf") + return ( + (b.max_time_seconds - b.used_time_seconds) + if b.max_time_seconds is not None + else float("inf") + ) return float("inf") # ========== Tool Methods ========== diff --git a/tests/unit/test_context_compression.py b/tests/unit/test_context_compression.py index 73c626b8..b0206476 100644 --- a/tests/unit/test_context_compression.py +++ b/tests/unit/test_context_compression.py @@ -192,6 +192,18 @@ def test_compress_context_dedup_preserves_distinct_tool_call_payloads() -> None: assert tool_result_ids == ["tc_1", "tc_2"] +def test_compress_context_dedup_handles_unhashable_payload_values() -> None: + ctx = Context(config=Config()) + ctx.stm_add(Message(role="assistant", text="tool result", data={"values": {1, 2, 3}})) + ctx.stm_add(Message(role="assistant", text="tool result", data={"values": {3, 2, 1}})) + + compress_context(ctx, strategy="dedup_then_truncate", max_messages=10) + + messages = ctx.stm_get() + assert len(messages) == 1 + assert messages[0].data == {"values": {1, 2, 3}} + + def test_compress_context_target_tokens_trims_long_history() -> None: ctx = Context(config=Config()) for idx in range(8): diff --git a/tests/unit/test_context_implementation.py b/tests/unit/test_context_implementation.py index 2ae37a08..43488ef4 100644 --- a/tests/unit/test_context_implementation.py +++ b/tests/unit/test_context_implementation.py @@ -201,6 +201,26 @@ def test_context_assemble_degrades_when_token_budget_low(): assert assembled.metadata["retrieval"]["degrade_reason"] == "token_budget_low" +def test_context_assemble_zero_token_budget_skips_retrieval() -> None: + ltm = _FakeRetrieval([Message(role="assistant", text="ltm-hit")]) + knowledge = _FakeRetrieval([Message(role="assistant", text="knowledge-hit")]) + ctx = Context( + config=Config(), + budget=Budget(max_tokens=0), + long_term_memory=ltm, + knowledge=knowledge, + ) + ctx.stm_add(Message(role="user", text="query")) + + assembled = ctx.assemble() + + assert [message.text for message in assembled.messages] == ["query"] + assert assembled.metadata["retrieval"]["degraded"] is True + assert assembled.metadata["retrieval"]["degrade_reason"] == "token_budget_low" + assert ltm.calls == [] + assert knowledge.calls == [] + + def test_context_assemble_handles_retrieval_exception_gracefully(): ltm = _FakeRetrieval([Message(role="assistant", text="ltm-hit")], fail=True) knowledge = _FakeRetrieval([Message(role="assistant", text="knowledge-hit")]) From e5281484559833f46aecf25e61ba58798f7786e7 Mon Sep 17 00:00:00 2001 From: mindfn Date: Tue, 10 Mar 2026 16:12:17 +0800 Subject: [PATCH 5/5] Narrow PR 211 to test and example alignment Drop the production-side compression and auto-compress changes from PR 211 and keep the branch focused on correcting outdated expectations in tests and examples. Key changes: - remove the framework-level compression.core contract from unit coverage - rewrite Context tests around the current default assemble path and moving-compressor behavior - update ReactAgent tests to use the existing Message-level execute contract and remove auto_compress expectations - refresh the AgentScope compat example docs/comments to describe moving compression as the current framework behavior Rationale: The module owner confirmed the failing assumptions were in UT/example coverage rather than in mainline runtime behavior, so this PR should only realign tests and example documentation with the current implementation. --- dare_framework/agent/react_agent.py | 105 +--- dare_framework/compression/__init__.py | 9 +- dare_framework/compression/core.py | 472 ------------------ dare_framework/context/context.py | 252 +--------- dare_framework/context/kernel.py | 4 +- dare_framework/memory/in_memory_stm.py | 8 +- .../DESIGN.md | 10 +- .../README.md | 2 +- .../compat_agent.py | 9 +- tests/unit/test_context_compression.py | 313 ------------ tests/unit/test_context_implementation.py | 466 ++--------------- .../unit/test_example_10_agentscope_compat.py | 2 + .../test_react_agent_gateway_injection.py | 181 +------ 13 files changed, 105 insertions(+), 1728 deletions(-) delete mode 100644 dare_framework/compression/core.py delete mode 100644 tests/unit/test_context_compression.py diff --git a/dare_framework/agent/react_agent.py b/dare_framework/agent/react_agent.py index 0a28f568..b6367958 100644 --- a/dare_framework/agent/react_agent.py +++ b/dare_framework/agent/react_agent.py @@ -112,11 +112,6 @@ def __init__( tool_gateway: IToolGateway, plan_provider: IToolProvider | None = None, max_tool_rounds: int = 10, - auto_compress: bool = False, - compress_trigger_ratio: float = 0.9, - compress_target_ratio: float = 0.75, - compress_max_messages: int | None = None, - compress_strategy: str = "dedup_then_truncate", agent_channel: AgentChannel | None = None, ) -> None: super().__init__(name, agent_channel=agent_channel) @@ -125,19 +120,6 @@ def __init__( self._context = context self._tool_gateway = tool_gateway self._plan_provider = plan_provider - self._auto_compress = bool(auto_compress) - self._compress_trigger_ratio = _clamp_ratio(compress_trigger_ratio, default=0.9) - self._compress_target_ratio = _clamp_ratio(compress_target_ratio, default=0.75) - self._compress_max_messages = ( - compress_max_messages - if isinstance(compress_max_messages, int) and compress_max_messages > 0 - else None - ) - self._compress_strategy = ( - compress_strategy.strip() - if isinstance(compress_strategy, str) and compress_strategy.strip() - else "dedup_then_truncate" - ) self._context.set_tool_gateway(self._tool_gateway) # 运行时检测是否可以启用 SmartContext 能力 @@ -154,7 +136,7 @@ def plan_provider(self) -> IToolProvider | None: async def execute( self, - task: Message | str, + task: Message, *, transport: AgentChannel | None = None, ) -> RunResult: @@ -165,13 +147,12 @@ async def execute( async def _execute_basic( self, - task: Message | str, + task: Message, *, transport: AgentChannel | None = None, ) -> RunResult: """原始基础 ReAct 循环实现。""" - user_message = task if isinstance(task, Message) else Message(role="user", text=task) - self._context.stm_add(user_message) + self._context.stm_add(task) gateway = self._tool_gateway @@ -183,9 +164,6 @@ async def _execute_basic( print(f"[{self.name}] Round {round_idx + 1}/{self._max_tool_rounds}: 调用模型中...", flush=True) assembled = await self._context.assemble_for_model() messages = self._build_model_messages(assembled) - if self._maybe_auto_compress(messages): - assembled = await self._context.assemble_for_model() - messages = self._build_model_messages(assembled) model_input = ModelInput( messages=messages, @@ -351,7 +329,7 @@ async def _execute_basic( async def _execute_with_smart_context( self, - task: Message | str, + task: Message, *, transport: AgentChannel | None = None, ) -> RunResult: @@ -366,7 +344,7 @@ async def _execute_with_smart_context( return await self._execute_basic(task, transport=transport) _ = transport - source_user_message = task if isinstance(task, Message) else Message(role="user", text=task) + source_user_message = task user_message = Message( role=source_user_message.role, kind=source_user_message.kind, @@ -427,26 +405,6 @@ async def _execute_with_smart_context( messages.append(self._next_round_reflection_prompt) self._next_round_reflection_prompt = None - if self._maybe_auto_compress(messages): - assembled = await self._context.assemble_for_model() - messages = list(assembled.messages) - prompt_def = getattr(assembled, "sys_prompt", None) - sys_prompt_message = ( - Message( - role=prompt_def.role, - text=prompt_def.content, - name=prompt_def.name, - metadata=dict(prompt_def.metadata), - mark=MessageMark.IMMUTABLE, - id="sys_prompt", - ) - if prompt_def is not None - else None - ) - messages = self._context.order_messages_for_llm(messages, sys_prompt_message) - if injected_reflection_prompt is not None: - messages.append(injected_reflection_prompt) - # Inject critical_block from plan_provider (maintained by plan tools) # Disabled: skip injection to observe plan agent behavior without it if False and self._plan_provider is not None: @@ -650,38 +608,6 @@ def _build_model_messages(self, assembled: Any) -> list[Message]: ) return messages - def _maybe_auto_compress(self, model_messages: list[Message]) -> bool: - """Auto-compress context before model invocation when token estimate is near budget.""" - if not self._auto_compress: - return False - max_tokens = self._context.budget.max_tokens - if max_tokens is None or max_tokens <= 0: - return False - - estimated_tokens = _estimate_messages_tokens(model_messages) - trigger_tokens = max(1, int(max_tokens * self._compress_trigger_ratio)) - if estimated_tokens < trigger_tokens: - return False - - stm_messages = self._context.stm_get() - if not stm_messages: - return False - - max_messages = self._compress_max_messages - if max_messages is None: - max_messages = max(1, int(len(stm_messages) * self._compress_target_ratio)) - if max_messages >= len(stm_messages): - max_messages = max(1, len(stm_messages) - 1) - - target_tokens = max(1, int(max_tokens * self._compress_target_ratio)) - self._context.compress( - strategy=self._compress_strategy, - max_messages=max_messages, - target_tokens=target_tokens, - tool_pair_safe=True, - ) - return True - async def _emit_terminal_transport_message( self, *, @@ -823,25 +749,4 @@ def _tool_calls_signature(tool_calls: list[dict[str, Any]]) -> tuple[str, ...]: return tuple(signature) -def _estimate_messages_tokens(messages: list[Message]) -> int: - total = 0 - for message in messages: - content = (message.text or "").strip() - attachment_tokens = len(message.attachments) * 32 - total += max(1, len(content) // 4) + attachment_tokens + 8 - return total - - -def _clamp_ratio(value: Any, *, default: float) -> float: - try: - ratio = float(value) - except (TypeError, ValueError): - return default - if not math.isfinite(ratio) or ratio <= 0: - return default - if ratio > 1: - return 1.0 - return ratio - - __all__ = ["ReactAgent"] diff --git a/dare_framework/compression/__init__.py b/dare_framework/compression/__init__.py index c8020700..e0b77334 100644 --- a/dare_framework/compression/__init__.py +++ b/dare_framework/compression/__init__.py @@ -1,8 +1,11 @@ -"""Compression utilities for context and memories.""" +"""Compression utilities for context and memories. + +- MovingCompressor: 移动窗口式 STM 压缩(LLM 摘要),见 moving_compression。 +""" from __future__ import annotations -from .core import compress_context, compress_context_llm_summary from .moving_compression import MovingCompressor -__all__ = ["compress_context", "compress_context_llm_summary", "MovingCompressor"] +__all__ = ["MovingCompressor"] + diff --git a/dare_framework/compression/core.py b/dare_framework/compression/core.py deleted file mode 100644 index 2273eef2..00000000 --- a/dare_framework/compression/core.py +++ /dev/null @@ -1,472 +0,0 @@ -"""Core context compression helpers. - -This module preserves the synchronous compression entrypoints described by the -design docs while the moving-window compressor remains available for -`assemble_for_model()` flows. -""" - -from __future__ import annotations - -from dataclasses import asdict, is_dataclass -from typing import TYPE_CHECKING, Any, List, Tuple - -from dare_framework.context.types import Message as CtxMessage, MessageKind, MessageMark -from dare_framework.model import ModelInput - -if TYPE_CHECKING: - from dare_framework.context.kernel import IContext - from dare_framework.context.types import Message - from dare_framework.model import IModelAdapter - - -_UNCHANGED = object() - - -def _freeze_value(value: Any) -> Any: - """Build a hashable structural key for nested message payloads.""" - if isinstance(value, dict): - return tuple(sorted((str(key), _freeze_value(item)) for key, item in value.items())) - if isinstance(value, list): - return tuple(_freeze_value(item) for item in value) - if isinstance(value, tuple): - return tuple(_freeze_value(item) for item in value) - if isinstance(value, (set, frozenset)): - frozen_items = [_freeze_value(item) for item in value] - return tuple(sorted(frozen_items, key=repr)) - if is_dataclass(value) and not isinstance(value, type): - return ("dataclass", type(value).__qualname__, _freeze_value(asdict(value))) - if hasattr(value, "kind") and hasattr(value, "uri"): - return ( - getattr(value.kind, "value", value.kind), - value.uri, - value.mime_type, - value.filename, - _freeze_value(getattr(value, "metadata", {})), - ) - try: - hash(value) - except TypeError: - if hasattr(value, "__dict__"): - return (type(value).__qualname__, _freeze_value(vars(value))) - return (type(value).__qualname__, repr(value)) - return value - - -def _copy_message( - message: Message, - *, - data: dict[str, Any] | None | object = _UNCHANGED, - metadata: dict[str, Any] | object = _UNCHANGED, -) -> Message: - return CtxMessage( - role=message.role, - kind=message.kind, - text=message.text, - attachments=list(message.attachments), - data=message.data if data is _UNCHANGED else data, - name=message.name, - metadata=message.metadata if metadata is _UNCHANGED else dict(metadata), - mark=getattr(message, "mark", MessageMark.TEMPORARY), - id=getattr(message, "id", None), - ) - - -def _dedup_messages(messages: List[Message]) -> Tuple[List[Message], int]: - """De-duplicate only when the full public message payload matches.""" - seen: set[Any] = set() - result: List[Message] = [] - removed = 0 - - for msg in messages: - key = ( - msg.role, - msg.kind, - msg.text, - msg.name, - _freeze_value(msg.attachments), - _freeze_value(msg.data), - ) - if key in seen: - removed += 1 - continue - seen.add(key) - result.append(msg) - - return result, removed - - -def _build_summary_preview( - messages: List[Message], - max_messages: int, - tail_max: int = 10, -) -> Tuple[List[Message], int]: - """Heuristic, model-free summary strategy.""" - total = len(messages) - if total <= max_messages or max_messages <= 1: - return messages, 0 - - tail_capacity = max_messages - 1 - keep_tail = min(tail_max, tail_capacity, total) - head = messages[:-keep_tail] if keep_tail > 0 else messages - tail = messages[-keep_tail:] if keep_tail > 0 else [] - - if not head: - return messages, 0 - - preview_lines: List[str] = [] - for msg in head: - content = (msg.text or "").strip() - if not content: - continue - snippet = content.replace("\n", " ") - if len(snippet) > 120: - snippet = snippet[:120] + "..." - preview_lines.append(f"{msg.role}: {snippet}") - - if not preview_lines: - return messages, 0 - - summary_text = "Conversation summary (heuristic, no LLM):\n" + "\n".join(preview_lines) - summary_message = CtxMessage( - role="system", - kind=MessageKind.SUMMARY, - text=summary_text, - metadata={"compressed": True, "strategy": "summary_preview"}, - ) - - new_messages: List[Message] = [summary_message, *tail] - removed = total - len(new_messages) - return new_messages, removed - - -def _estimate_tokens(messages: List[Message]) -> int: - """Rough token estimate using a cheap character + attachment heuristic.""" - total = 0 - for msg in messages: - content = (msg.text or "").strip() - attachment_tokens = len(msg.attachments) * 32 - total += max(1, len(content) // 4) + attachment_tokens + 8 - return total - - -def _trim_to_target_tokens(messages: List[Message], target_tokens: int | None) -> Tuple[List[Message], int]: - """Trim oldest messages until estimated token size fits target_tokens.""" - if target_tokens is None or target_tokens <= 0: - return messages, 0 - if _estimate_tokens(messages) <= target_tokens: - return messages, 0 - - trimmed = list(messages) - removed = 0 - while len(trimmed) > 1 and _estimate_tokens(trimmed) > target_tokens: - removable_idx = next( - ( - idx - for idx, message in enumerate(trimmed) - if getattr(message, "mark", MessageMark.TEMPORARY) - not in (MessageMark.IMMUTABLE, MessageMark.PERSISTENT) - ), - None, - ) - if removable_idx is None: - break - trimmed.pop(removable_idx) - removed += 1 - return trimmed, removed - - -def _extract_tool_call_ids(message: Message) -> list[str]: - """Collect tool call ids declared on an assistant message.""" - if message.role != "assistant": - return [] - tool_calls = message.data.get("tool_calls", []) if isinstance(message.data, dict) else [] - if not isinstance(tool_calls, list): - return [] - - ids: list[str] = [] - for call in tool_calls: - if not isinstance(call, dict): - continue - tool_id = call.get("id") - if isinstance(tool_id, str) and tool_id.strip(): - ids.append(tool_id.strip()) - return ids - - -def _enforce_tool_pair_safety(messages: List[Message]) -> Tuple[List[Message], int]: - """Keep tool_call/tool_result in sync so compression never leaves orphan pairs.""" - tool_result_ids = { - message.name.strip() - for message in messages - if message.role == "tool" and isinstance(message.name, str) and message.name.strip() - } - - updated_messages: list[Message] = [] - retained_call_ids: set[str] = set() - retained_idless_tool_names: set[str] = set() - changes = 0 - - for message in messages: - if message.role != "assistant": - updated_messages.append(message) - continue - raw_calls = message.data.get("tool_calls", []) if isinstance(message.data, dict) else [] - if not isinstance(raw_calls, list): - updated_messages.append(message) - continue - - filtered_calls = [] - for call in raw_calls: - if not isinstance(call, dict): - continue - tool_id = call.get("id") - if isinstance(tool_id, str) and tool_id.strip() and tool_id.strip() in tool_result_ids: - filtered_calls.append(call) - retained_call_ids.add(tool_id.strip()) - continue - if not isinstance(tool_id, str) or not tool_id.strip(): - filtered_calls.append(call) - tool_name = call.get("name") - if isinstance(tool_name, str) and tool_name.strip(): - retained_idless_tool_names.add(tool_name.strip()) - - if len(filtered_calls) != len(raw_calls): - changes += len(raw_calls) - len(filtered_calls) - updated_data = dict(message.data) if isinstance(message.data, dict) else {} - updated_data["tool_calls"] = filtered_calls - updated_messages.append(_copy_message(message, data=updated_data)) - else: - retained_call_ids.update(_extract_tool_call_ids(message)) - updated_messages.append(message) - - final_messages: list[Message] = [] - for message in updated_messages: - if message.role == "tool": - tool_id = message.name.strip() if isinstance(message.name, str) else "" - if tool_id in retained_call_ids or (tool_id and tool_id in retained_idless_tool_names): - final_messages.append(message) - continue - if tool_id: - changes += 1 - continue - final_messages.append(message) - return final_messages, changes - - -def _annotate_strategy(messages: List[Message], strategy: str) -> List[Message]: - """Attach strategy metadata to the first message when compression rewrites context.""" - if not messages: - return messages - for message in messages: - if message.metadata.get("compressed") is True: - return messages - - head = messages[0] - metadata = dict(head.metadata) - metadata["compressed"] = True - metadata.setdefault("strategy", strategy) - messages[0] = _copy_message(head, metadata=metadata) - return messages - - -def compress_context( - context: IContext, - *, - phase: str | None = None, - max_messages: int | None = None, - **options: Any, -) -> None: - """Compress short-term memory for a given context.""" - _ = phase - - target_tokens_raw = options.get("target_tokens") - target_tokens: int | None = None - if target_tokens_raw is not None: - try: - target_tokens = int(target_tokens_raw) - except (TypeError, ValueError): - target_tokens = None - - if (max_messages is None or max_messages < 0) and (target_tokens is None or target_tokens <= 0): - return - - stm_get = getattr(context, "stm_get", None) - stm_clear = getattr(context, "stm_clear", None) - stm_add = getattr(context, "stm_add", None) - if not callable(stm_get) or not callable(stm_clear) or not callable(stm_add): - return - - messages: List[Message] = list(stm_get()) - if not messages: - return - - if max_messages is None: - max_messages = len(messages) - elif max_messages < 0: - max_messages = len(messages) - - strategy = options.get("strategy", "truncate") - tool_pair_safe = bool(options.get("tool_pair_safe", False)) - - removed_total = 0 - - if strategy == "summary_preview": - messages, removed = _build_summary_preview(messages, max_messages) - removed_total += removed - - if strategy == "dedup_then_truncate": - messages, removed = _dedup_messages(messages) - removed_total += removed - - if max_messages == 0: - removed_total += len(messages) - messages = [] - elif len(messages) > max_messages: - protected = [ - message - for message in messages - if getattr(message, "mark", MessageMark.TEMPORARY) - in (MessageMark.IMMUTABLE, MessageMark.PERSISTENT) - ] - temporary = [ - message - for message in messages - if getattr(message, "mark", MessageMark.TEMPORARY) - not in (MessageMark.IMMUTABLE, MessageMark.PERSISTENT) - ] - keep_temporary = max(max_messages - len(protected), 0) - if keep_temporary <= 0: - kept_tail: list[Message] = [] - elif keep_temporary < len(temporary): - kept_tail = temporary[-keep_temporary:] - else: - kept_tail = temporary - kept_tail_refs = {id(message) for message in kept_tail} - messages = [ - message - for message in messages - if ( - getattr(message, "mark", MessageMark.TEMPORARY) - in (MessageMark.IMMUTABLE, MessageMark.PERSISTENT) - ) - or id(message) in kept_tail_refs - ] - removed_total += len(protected) + len(temporary) - len(messages) - - messages, removed = _trim_to_target_tokens(messages, target_tokens) - removed_total += removed - - if tool_pair_safe: - messages, changes = _enforce_tool_pair_safety(messages) - removed_total += changes - - if removed_total == 0: - return - - messages = _annotate_strategy(messages, str(strategy)) - stm_clear() - for msg in messages: - stm_add(msg) - - -async def compress_context_llm_summary( - context: IContext, - *, - model: IModelAdapter, - max_messages: int, - keep_tail: int = 8, - system_prompt: str | None = None, - language: str = "zh", -) -> None: - """High-level compression using the LLM to generate a semantic summary.""" - if max_messages <= 1: - return - - stm_get = getattr(context, "stm_get", None) - stm_clear = getattr(context, "stm_clear", None) - stm_add = getattr(context, "stm_add", None) - if not callable(stm_get) or not callable(stm_clear) or not callable(stm_add): - return - - messages: List[Message] = list(stm_get()) - total = len(messages) - if total <= max_messages: - return - - tail_capacity = max_messages - 1 - keep_tail_eff = min(keep_tail, tail_capacity, total) - head = messages[:-keep_tail_eff] if keep_tail_eff > 0 else messages - tail = messages[-keep_tail_eff:] if keep_tail_eff > 0 else [] - - if not head: - return - - lines: List[str] = [] - for msg in head: - content = (msg.text or "").strip() - if not content: - continue - snippet = content.replace("\n", " ") - if len(snippet) > 512: - snippet = snippet[:512] + "..." - lines.append(f"{msg.role}: {snippet}") - - if not lines: - return - - conversation_text = "\n".join(lines) - if system_prompt is None: - if language == "zh": - system_prompt = ( - "你是一个对话摘要助手,请在不丢失关键信息的前提下," - "用简洁、结构化的方式总结下面的一段历史对话。" - "可以合并重复信息,但不要编造不存在的内容。" - ) - else: - system_prompt = ( - "You are a conversation summarization assistant. " - "Produce a concise, structured summary of the following history, " - "preserving key facts and decisions. Do not invent new information." - ) - - sys_msg = CtxMessage(role="system", kind=MessageKind.SUMMARY, text=system_prompt) - user_intro = ( - "下面是一段需要被压缩的历史对话,请输出一个摘要,用于后续继续对话使用。\n\n" - "=== 历史开始 ===\n" - f"{conversation_text}\n" - "=== 历史结束 ===" - if language == "zh" - else - "Here is the conversation history that needs to be compressed. " - "Please output a summary that can be used for continuing the dialogue.\n\n" - "=== HISTORY START ===\n" - f"{conversation_text}\n" - "=== HISTORY END ===" - ) - user_msg = CtxMessage(role="user", text=user_intro) - - model_input = ModelInput( - messages=[sys_msg, user_msg], - tools=[], - metadata={"compression": "llm_summary"}, - ) - - response = await model.generate(model_input) - summary_text = (response.content or "").strip() - if not summary_text: - return - - summary_message = CtxMessage( - role="system", - kind=MessageKind.SUMMARY, - text=summary_text, - metadata={"compressed": True, "strategy": "llm_summary"}, - ) - - stm_clear() - stm_add(summary_message) - for msg in tail: - stm_add(msg) - - -__all__ = ["compress_context", "compress_context_llm_summary"] diff --git a/dare_framework/context/context.py b/dare_framework/context/context.py index a6b5ebbf..70b37cae 100644 --- a/dare_framework/context/context.py +++ b/dare_framework/context/context.py @@ -176,17 +176,13 @@ def budget_remaining(self, resource: str) -> float: """Get remaining budget for a resource.""" b = self._budget if resource == "tokens": - return (b.max_tokens - b.used_tokens) if b.max_tokens is not None else float("inf") + return (b.max_tokens - b.used_tokens) if b.max_tokens else float("inf") elif resource == "cost": - return (b.max_cost - b.used_cost) if b.max_cost is not None else float("inf") + return (b.max_cost - b.used_cost) if b.max_cost else float("inf") elif resource == "tool_calls": - return (b.max_tool_calls - b.used_tool_calls) if b.max_tool_calls is not None else float("inf") + return (b.max_tool_calls - b.used_tool_calls) if b.max_tool_calls else float("inf") elif resource == "time_seconds": - return ( - (b.max_time_seconds - b.used_time_seconds) - if b.max_time_seconds is not None - else float("inf") - ) + return (b.max_time_seconds - b.used_time_seconds) if b.max_time_seconds else float("inf") return float("inf") # ========== Tool Methods ========== @@ -205,52 +201,21 @@ def list_tools(self) -> list[CapabilityDescriptor]: def assemble(self) -> AssembledContext: return self._assemble_context.assemble(self) - def compress(self, **options: Any) -> None: - """Compress context to fit within budget.""" - from dare_framework.compression.core import compress_context - - # Preserve backend STM semantics (for example SmartSTM mark-based retention) - # before applying advanced compression strategies. - compress_impl = getattr(self._short_term_memory, "compress", None) - has_advanced_options = any( - key in options - for key in ("target_tokens", "tool_pair_safe", "strategy", "phase") - ) - raw_max_messages = options.get("max_messages") - max_messages = ( - raw_max_messages - if isinstance(raw_max_messages, int) and raw_max_messages >= 0 - else None - ) - if callable(compress_impl) and not has_advanced_options: - compress_impl(max_messages=max_messages) + async def compress(self, **options: Any) -> None: + """压缩 STM:仅委托 moving_compressor.prune;无 compressor 时无操作。""" + if self._moving_compressor is None: return - - compress_context(self, **options) - - # For advanced compression, run backend max-message retention after strategy - # execution so strategy implementations can inspect full pre-trim history. - if callable(compress_impl) and has_advanced_options and max_messages is not None: - compress_impl(max_messages=max_messages) + # 只把与 token 预算相关的参数传给压缩器;摘要 prompt 与语言策略由压缩器内部自行决定。 + prune_opts: dict[str, Any] = {} + if "max_context_tokens" in options: + prune_opts["max_context_tokens"] = options["max_context_tokens"] + elif self._context_window_tokens is not None and self._context_window_tokens > 0: + prune_opts["max_context_tokens"] = self._context_window_tokens + await self._moving_compressor.prune(self, **prune_opts) async def assemble_for_model(self, **options: Any) -> AssembledContext: - """Model-facing assembly path with optional moving-window compression.""" - sync_compress_keys = ("max_messages", "target_tokens", "tool_pair_safe", "strategy", "phase") - has_sync_compress_options = any(key in options for key in sync_compress_keys) - - # Avoid calling compress({}) on the default assembly path. Some callers - # override compress() for observability, and an empty no-op call changes - # behavior without providing any trimming value. - if has_sync_compress_options: - self.compress(**options) - - if self._moving_compressor is not None: - prune_opts: dict[str, Any] = {} - if "max_context_tokens" in options: - prune_opts["max_context_tokens"] = options["max_context_tokens"] - elif self._context_window_tokens is not None and self._context_window_tokens > 0: - prune_opts["max_context_tokens"] = self._context_window_tokens - await self._moving_compressor.prune(self, **prune_opts) + """供模型调用的装配入口:在内部静默触发压缩,然后返回 AssembledContext。""" + await self.compress(**options) return self.assemble() @@ -258,81 +223,6 @@ class DefaultAssembledContext(IAssembleContext): """Default context assembly strategy. """ - _DEFAULT_TOP_K = 3 - _DEFAULT_RESERVE_TOKENS = 256 - _DEFAULT_SOURCE_RATIO = 0.5 - - def _safe_int(self, value: Any, default: int, *, minimum: int | None = None) -> int: - try: - parsed = int(value) - except (TypeError, ValueError, OverflowError): - parsed = default - if minimum is not None: - parsed = max(minimum, parsed) - return parsed - - def _safe_ratio(self, value: Any) -> float: - try: - parsed = float(value) - except (TypeError, ValueError, OverflowError): - return self._DEFAULT_SOURCE_RATIO - if not math.isfinite(parsed) or parsed < 0: - return self._DEFAULT_SOURCE_RATIO - return parsed - - def _estimate_tokens(self, messages: list[Message]) -> int: - """Rough token estimate using character + attachment heuristics.""" - total = 0 - for message in messages: - content_tokens = max(1, len((message.text or "").strip()) // 4) - attachment_tokens = len(message.attachments) * 32 - total += content_tokens + attachment_tokens + 8 - return total - - def _take_with_budget(self, messages: list[Message], budget_tokens: float) -> list[Message]: - if budget_tokens == float("inf"): - return list(messages) - if budget_tokens <= 0: - return [] - - kept: list[Message] = [] - used = 0 - for message in messages: - message_tokens = self._estimate_tokens([message]) - if used + message_tokens > budget_tokens: - # Skip oversized candidates so smaller later hits can still fit. - continue - kept.append(message) - used += message_tokens - return kept - - def _derive_query(self, messages: list[Message]) -> str: - for message in reversed(messages): - if message.role == "user" and (message.text or "").strip(): - return (message.text or "").strip() - for message in reversed(messages): - if (message.text or "").strip(): - return (message.text or "").strip() - return "" - - def _load_source_options(self, config_map: dict[str, Any]) -> tuple[int, float]: - top_k = self._safe_int( - config_map.get("assemble_top_k"), - self._DEFAULT_TOP_K, - minimum=0, - ) - ratio = self._safe_ratio(config_map.get("assemble_ratio", self._DEFAULT_SOURCE_RATIO)) - return top_k, ratio - - def _set_degrade( - self, - retrieval_metadata: dict[str, Any], - *, - reason: str, - ) -> None: - retrieval_metadata["degraded"] = True - retrieval_metadata["degrade_reason"] = reason - def assemble(self, context: IContext) -> AssembledContext: messages = context.stm_get() tools = context.list_tools() @@ -342,120 +232,12 @@ def assemble(self, context: IContext) -> AssembledContext: sys_prompt = enrich_prompt_with_skill(sys_prompt, context.sys_skill) - query = self._derive_query(messages) - ltm_config = context.config.long_term_memory if isinstance(context.config.long_term_memory, dict) else {} - knowledge_config = context.config.knowledge if isinstance(context.config.knowledge, dict) else {} - ltm_top_k, ltm_ratio = self._load_source_options(ltm_config) - knowledge_top_k, knowledge_ratio = self._load_source_options(knowledge_config) - ltm_active = context.long_term_memory is not None and ltm_top_k > 0 - knowledge_active = context.knowledge is not None and knowledge_top_k > 0 - - if ltm_active and not knowledge_active: - reserve_tokens_raw = ltm_config.get("assemble_reserve_tokens") - elif knowledge_active and not ltm_active: - reserve_tokens_raw = knowledge_config.get("assemble_reserve_tokens") - else: - reserve_tokens_raw = ltm_config.get("assemble_reserve_tokens") - if reserve_tokens_raw is None: - reserve_tokens_raw = knowledge_config.get("assemble_reserve_tokens") - reserve_tokens = self._safe_int( - reserve_tokens_raw, - self._DEFAULT_RESERVE_TOKENS, - minimum=0, - ) - - retrieval_metadata: dict[str, Any] = { - "query": query, - "stm_count": len(messages), - "ltm_requested": ltm_top_k, - "knowledge_requested": knowledge_top_k, - "ltm_count": 0, - "knowledge_count": 0, - "degraded": False, - "degrade_reason": None, - } - - remaining_tokens = context.budget_remaining("tokens") - stm_token_estimate = self._estimate_tokens(messages) - retrieval_budget: float = float("inf") - if remaining_tokens != float("inf"): - retrieval_budget = max(0.0, float(remaining_tokens) - float(stm_token_estimate) - float(reserve_tokens)) - - ltm_messages: list[Message] = [] - knowledge_messages: list[Message] = [] - - if retrieval_budget <= 0 and (ltm_active or knowledge_active): - self._set_degrade(retrieval_metadata, reason="token_budget_low") - else: - ratio_total = 0.0 - if ltm_active: - ratio_total += ltm_ratio - if knowledge_active: - ratio_total += knowledge_ratio - if ratio_total <= 0: - # Fall back only across active retrieval sources. - ratio_total = 0.0 - if ltm_active: - ltm_ratio = self._DEFAULT_SOURCE_RATIO - ratio_total += ltm_ratio - if knowledge_active: - knowledge_ratio = self._DEFAULT_SOURCE_RATIO - ratio_total += knowledge_ratio - - normalized_ltm_ratio = (ltm_ratio / ratio_total) if ltm_active and ratio_total > 0 else 0.0 - normalized_knowledge_ratio = ( - (knowledge_ratio / ratio_total) if knowledge_active and ratio_total > 0 else 0.0 - ) - - ltm_budget = float("inf") - knowledge_budget = float("inf") - if retrieval_budget != float("inf"): - ltm_budget = retrieval_budget * normalized_ltm_ratio - knowledge_budget = retrieval_budget * normalized_knowledge_ratio - - ltm_retrieval_failed = False - if ltm_active: - if ltm_budget <= 0: - ltm_messages = [] - else: - try: - ltm_candidates = context.long_term_memory.get(query=query, top_k=ltm_top_k) - ltm_messages = self._take_with_budget(ltm_candidates, ltm_budget) - if len(ltm_messages) < len(ltm_candidates): - self._set_degrade(retrieval_metadata, reason="token_budget_low") - except Exception: - ltm_retrieval_failed = True - self._set_degrade(retrieval_metadata, reason="ltm_retrieval_failed") - ltm_messages = [] - - if knowledge_active: - try: - effective_knowledge_budget = knowledge_budget - if retrieval_budget != float("inf") and ltm_active and ltm_retrieval_failed: - effective_knowledge_budget = retrieval_budget - if effective_knowledge_budget <= 0: - knowledge_messages = [] - else: - knowledge_candidates = context.knowledge.get(query=query, top_k=knowledge_top_k) - knowledge_messages = self._take_with_budget(knowledge_candidates, effective_knowledge_budget) - if len(knowledge_messages) < len(knowledge_candidates): - self._set_degrade(retrieval_metadata, reason="token_budget_low") - except Exception: - if not retrieval_metadata["degraded"]: - self._set_degrade(retrieval_metadata, reason="knowledge_retrieval_failed") - knowledge_messages = [] - - merged_messages = [*messages, *ltm_messages, *knowledge_messages] - retrieval_metadata["ltm_count"] = len(ltm_messages) - retrieval_metadata["knowledge_count"] = len(knowledge_messages) - return AssembledContext( - messages=merged_messages, + messages=list(messages), sys_prompt=sys_prompt, tools=tools, metadata={ "context_id": context.id, - "retrieval": retrieval_metadata, }, ) diff --git a/dare_framework/context/kernel.py b/dare_framework/context/kernel.py index 25b2cc48..50193a25 100644 --- a/dare_framework/context/kernel.py +++ b/dare_framework/context/kernel.py @@ -104,9 +104,9 @@ def list_tools(self) -> list[CapabilityDescriptor]: ... def assemble(self) -> AssembledContext: ... - # Compress (core):同步高级压缩入口;assemble_for_model 可在内部追加异步 moving compression。 + # Compress (core):由具体 Context 实现决定何时触发;默认在 assemble_for_model 中静默调用。 - def compress(self, **options: Any) -> None: ... + async def compress(self, **options: Any) -> None: ... # Assemble for model: 默认直接调用 assemble,由具体实现决定是否在内部触发 compress。 diff --git a/dare_framework/memory/in_memory_stm.py b/dare_framework/memory/in_memory_stm.py index 8035d1f4..a5e95c62 100644 --- a/dare_framework/memory/in_memory_stm.py +++ b/dare_framework/memory/in_memory_stm.py @@ -50,13 +50,7 @@ def compress(self, max_messages: int | None = None, **kwargs) -> int: Returns: Number of messages removed. """ - if max_messages is None: - return 0 - if max_messages <= 0: - removed_count = len(self._messages) - self._messages = [] - return removed_count - if len(self._messages) <= max_messages: + if max_messages is None or len(self._messages) <= max_messages: return 0 removed_count = len(self._messages) - max_messages diff --git a/examples/10-agentscope-compat-single-agent/DESIGN.md b/examples/10-agentscope-compat-single-agent/DESIGN.md index b292693e..47af4126 100644 --- a/examples/10-agentscope-compat-single-agent/DESIGN.md +++ b/examples/10-agentscope-compat-single-agent/DESIGN.md @@ -48,7 +48,7 @@ PlanNoteBook / SubTask / TruncatedFormatterBase / Knowledge / HttpStatefulClient - Memory: `dare_framework/memory/in_memory_stm.py` - Knowledge: `dare_framework/knowledge/kernel.py` - Plan: `dare_framework/plan_v2/types.py`, `dare_framework/plan_v2/tools.py` -- Compression: `dare_framework/compression/core.py` +- Compression: `dare_framework/compression/moving_compression.py` - MCP: `dare_framework/mcp/client.py`, `dare_framework/mcp/transports/http.py` ## 4. 能力差异矩阵(详细版) @@ -60,7 +60,7 @@ PlanNoteBook / SubTask / TruncatedFormatterBase / Knowledge / HttpStatefulClient | 循环结构 | `_reasoning() → _acting() → repeat` | `assemble() → generate() → tool calls → repeat` | 等价 | | 最大迭代 | `max_iterations=20` | `max_tool_rounds=10` | 等价(值不同) | | 并行 tool 执行 | `parallel_tool_calls=True` → `asyncio.gather` | 仅串行 | **Gap-R1** | -| 自动内存压缩 | `_compress_memory_if_needed()` 每轮触发 | 无自动压缩 | **Gap-R5** | +| 自动内存压缩 | `_compress_memory_if_needed()` 每轮触发 | 仅有 moving compression,未接入 ReAct 自动触发 | **Gap-R5** | | 超时 fallback | 超 max_iterations 做 summarization | 返回"未收敛"文本 | Gap-R2 | | Plan 注入 | `plan_to_hint()` 生成 `` | `critical_block` 注入 | 接近等价 | | Hook 粒度 | pre/post_reasoning, pre/post_acting | session/milestone/plan/tool 级 | Gap-R4 | @@ -162,8 +162,8 @@ PlanNoteBook / SubTask / TruncatedFormatterBase / Knowledge / HttpStatefulClient | 维度 | AgentScope | DARE | 差距 | |------|-----------|------|------| | 截断单位 | Token 数 | 消息条数 | **Gap-F2** | -| Tool pair 安全 | 成对删除 | 无保护 | **Gap-F1** | -| 自动触发 | 每次 `_reasoning()` 前 | 手动调用 | **Gap-F4** | +| Tool pair 安全 | 成对删除 | framework 无 formatter 级保护,由 Example 兼容层补齐 | **Gap-F1** | +| 自动触发 | 每次 `_reasoning()` 前 | moving compression 需显式接线,无 formatter 自动触发 | **Gap-F4** | | Provider 格式化 | OpenAI/Anthropic/Gemini/... formatter 子类 | 无 | Gap-F3 | | 标签感知 | 跳过 important 消息 | 无标签概念 | 依赖 Gap-M1 | @@ -218,7 +218,7 @@ PlanNoteBook / SubTask / TruncatedFormatterBase / Knowledge / HttpStatefulClient ### P1(高优先) - **Gap-M1**: Message 无 tag/mark - **Gap-Mem1/2/5**: InMemorySTM 无 mark/summary/tool-pair-safe compress -- **Gap-F1/F4**: compress_context 无 tool pair 安全/无自动触发 +- **Gap-F1/F4**: framework 仅提供 moving compression;无 formatter 级 tool pair 安全/无自动触发 - **Gap-R5**: ReactAgent 无自动内存压缩 - **Gap-LM4**: Usage 不规范化 reasoning_tokens - **Gap-S1/S2**: 无 StateModule/ISessionStore diff --git a/examples/10-agentscope-compat-single-agent/README.md b/examples/10-agentscope-compat-single-agent/README.md index ae901d8d..0bfb98f4 100644 --- a/examples/10-agentscope-compat-single-agent/README.md +++ b/examples/10-agentscope-compat-single-agent/README.md @@ -14,7 +14,7 @@ | 6 | `ChatModelBase` | `CompatFormattedModelAdapter` | E0/E2 | **Gap-LM1(P0)**(thinking), Gap-LM2(stream) | | 7 | `PlanNoteBook` | `CompatPlanNotebook` + 6 tools | E1 | Gap-P1(status), Gap-P5(序列化) | | 8 | `SubTask` | `CompatSubTask` | E1 | 合并于 Gap-P1 | -| 9 | `TruncatedFormatterBase` | `CompatTruncatedFormatter` | E1 | **Gap-F1**(tool pair safe), Gap-F2(token) | +| 9 | `TruncatedFormatterBase` | `CompatTruncatedFormatter` | E1 | **Gap-F1**(example 层补齐 tool pair safe), Gap-F2(token) | | 10 | `Knowledge` | `create_knowledge(rawdata)` | E0 | Gap-K1(embedding adapter) | | 11 | `HttpStatefulClient` | `HttpStatefulClientShim` | E1 | Gap-H3(缓存) | | 12 | `Session` | `JsonSessionBridge` | E2 | **Gap-S1**(StateModule), **Gap-S2**(ISessionStore) | diff --git a/examples/10-agentscope-compat-single-agent/compat_agent.py b/examples/10-agentscope-compat-single-agent/compat_agent.py index 0ad6bbee..6e1825b0 100644 --- a/examples/10-agentscope-compat-single-agent/compat_agent.py +++ b/examples/10-agentscope-compat-single-agent/compat_agent.py @@ -221,7 +221,7 @@ def to_framework_message(self) -> Message: # =========================================================================== # Capability 9: TruncatedFormatterBase — 截断格式化器 # AgentScope: token 级截断 + tool pair 安全 + provider-specific 格式化 -# DARE: compress_context() 按消息条数截断,无 tool pair 安全 [Gap-F1] +# DARE: 仅提供 moving compression;formatter 级 tool pair 安全需 Example 补齐 [Gap-F1] # =========================================================================== @@ -290,9 +290,10 @@ def _drop_with_tool_pair( ) -> int: """删除消息时保护 tool call/result 配对完整性。 - 这是 Gap-F1 的 Example 层补齐:框架的 compress_context() 不具备此能力。 - 当框架补齐 Gap-F1 (compress_context(tool_pair_safe=True)) 后, - 此方法应迁移到框架层。 + 这是 Gap-F1 的 Example 层补齐:当前框架只有 moving compression, + 并没有 formatter 级 tool pair 安全截断入口。 + 如果后续框架在公开 formatter/compression API 中补齐这层能力, + 此方法再考虑下沉到框架层。 """ removed_count = 0 removed_tool_ids: set[str] = set() diff --git a/tests/unit/test_context_compression.py b/tests/unit/test_context_compression.py deleted file mode 100644 index b0206476..00000000 --- a/tests/unit/test_context_compression.py +++ /dev/null @@ -1,313 +0,0 @@ -from __future__ import annotations - -from dare_framework.compression.core import compress_context -from dare_framework.config import Config -from dare_framework.context import AttachmentKind, AttachmentRef, Context, Message, MessageKind, MessageMark - - -def _tool_ids(message: Message) -> list[str]: - raw_calls = [] - if isinstance(message.data, dict): - raw_calls = message.data.get("tool_calls", []) - if not isinstance(raw_calls, list): - return [] - ids: list[str] = [] - for item in raw_calls: - if not isinstance(item, dict): - continue - tool_id = item.get("id") - if isinstance(tool_id, str) and tool_id: - ids.append(tool_id) - return ids - - -def test_compress_context_tool_pair_safe_removes_orphan_tool_result() -> None: - ctx = Context(config=Config()) - ctx.stm_add( - Message( - role="assistant", - text="tool call", - data={"tool_calls": [{"id": "tc_1", "name": "demo_tool", "arguments": {"x": 1}}]}, - ) - ) - ctx.stm_add(Message(role="tool", name="tc_1", text='{"success": true}')) - ctx.stm_add(Message(role="tool", name="tc_orphan", text='{"success": true}')) - - compress_context(ctx, strategy="truncate", max_messages=10, tool_pair_safe=True) - - messages = ctx.stm_get() - tool_names = [message.name for message in messages if message.role == "tool"] - assert "tc_1" in tool_names - assert "tc_orphan" not in tool_names - - -def test_compress_context_tool_pair_safe_removes_unmatched_tool_call_ids() -> None: - ctx = Context(config=Config()) - ctx.stm_add( - Message( - role="assistant", - text="tool call", - data={ - "tool_calls": [ - {"id": "tc_1", "name": "demo_tool", "arguments": {"x": 1}}, - {"id": "tc_2", "name": "missing_tool", "arguments": {"x": 2}}, - ] - }, - ) - ) - ctx.stm_add(Message(role="tool", name="tc_1", text='{"success": true}')) - - compress_context(ctx, strategy="truncate", max_messages=10, tool_pair_safe=True) - - assistant_message = next(message for message in ctx.stm_get() if message.role == "assistant") - assert _tool_ids(assistant_message) == ["tc_1"] - assert assistant_message.data == { - "tool_calls": [{"id": "tc_1", "name": "demo_tool", "arguments": {"x": 1}}] - } - - -def test_compress_context_tool_pair_safe_keeps_idless_tool_context() -> None: - ctx = Context(config=Config()) - ctx.stm_add( - Message( - role="assistant", - text="tool call without id", - data={"tool_calls": [{"name": "demo_tool", "arguments": {"x": 1}}]}, - ) - ) - ctx.stm_add(Message(role="tool", name="demo_tool", text='{"success": true}')) - - compress_context(ctx, strategy="truncate", max_messages=10, tool_pair_safe=True) - - messages = ctx.stm_get() - assistant_message = next(message for message in messages if message.role == "assistant") - raw_calls = assistant_message.data.get("tool_calls", []) if isinstance(assistant_message.data, dict) else [] - assert isinstance(raw_calls, list) - assert len(raw_calls) == 1 - assert any(message.role == "tool" and message.name == "demo_tool" for message in messages) - - -def test_compress_context_tool_pair_safe_drops_orphan_tool_results_with_mixed_id_modes() -> None: - ctx = Context(config=Config()) - ctx.stm_add( - Message( - role="assistant", - text="mixed tool calls", - data={ - "tool_calls": [ - {"id": "tc_1", "name": "demo_tool", "arguments": {"x": 1}}, - {"name": "demo_tool", "arguments": {"x": 2}}, - ] - }, - ) - ) - ctx.stm_add(Message(role="tool", name="tc_1", text='{"success": true}')) - ctx.stm_add(Message(role="tool", name="demo_tool", text='{"success": true}')) - ctx.stm_add(Message(role="tool", name="tc_orphan", text='{"success": true}')) - - compress_context(ctx, strategy="truncate", max_messages=10, tool_pair_safe=True) - - tool_names = [message.name for message in ctx.stm_get() if message.role == "tool"] - assert "tc_1" in tool_names - assert "demo_tool" in tool_names - assert "tc_orphan" not in tool_names - - -def test_compress_context_tool_pair_safe_preserves_assistant_id_and_mark_when_filtering_calls() -> None: - ctx = Context(config=Config()) - ctx.stm_add( - Message( - role="assistant", - text="mixed tool calls", - id="assistant-state", - mark=MessageMark.PERSISTENT, - data={ - "tool_calls": [ - {"id": "tc_1", "name": "demo_tool", "arguments": {"x": 1}}, - {"id": "tc_missing", "name": "demo_tool", "arguments": {"x": 2}}, - ] - }, - ) - ) - ctx.stm_add(Message(role="tool", name="tc_1", text='{"success": true}')) - - compress_context(ctx, strategy="truncate", max_messages=10, tool_pair_safe=True) - - assistant_message = next(message for message in ctx.stm_get() if message.role == "assistant") - assert assistant_message.id == "assistant-state" - assert assistant_message.mark == MessageMark.PERSISTENT - assert _tool_ids(assistant_message) == ["tc_1"] - assert assistant_message.data == { - "tool_calls": [{"id": "tc_1", "name": "demo_tool", "arguments": {"x": 1}}] - } - - -def test_compress_context_dedup_preserves_distinct_tool_call_payloads() -> None: - ctx = Context(config=Config()) - ctx.stm_add( - Message( - role="assistant", - kind=MessageKind.TOOL_CALL, - text="", - data={"tool_calls": [{"id": "tc_1", "name": "demo_tool", "arguments": {"x": 1}}]}, - ) - ) - ctx.stm_add( - Message( - role="tool", - kind=MessageKind.TOOL_RESULT, - name="tc_1", - text='{"success": true}', - data={"success": True}, - ) - ) - ctx.stm_add( - Message( - role="assistant", - kind=MessageKind.TOOL_CALL, - text="", - data={"tool_calls": [{"id": "tc_2", "name": "demo_tool", "arguments": {"x": 2}}]}, - ) - ) - ctx.stm_add( - Message( - role="tool", - kind=MessageKind.TOOL_RESULT, - name="tc_2", - text='{"success": true}', - data={"success": True}, - ) - ) - - compress_context(ctx, strategy="dedup_then_truncate", max_messages=10, tool_pair_safe=True) - - assistant_tool_ids = [ - _tool_ids(message) - for message in ctx.stm_get() - if message.role == "assistant" and message.kind == MessageKind.TOOL_CALL - ] - tool_result_ids = [message.name for message in ctx.stm_get() if message.role == "tool"] - - assert assistant_tool_ids == [["tc_1"], ["tc_2"]] - assert tool_result_ids == ["tc_1", "tc_2"] - - -def test_compress_context_dedup_handles_unhashable_payload_values() -> None: - ctx = Context(config=Config()) - ctx.stm_add(Message(role="assistant", text="tool result", data={"values": {1, 2, 3}})) - ctx.stm_add(Message(role="assistant", text="tool result", data={"values": {3, 2, 1}})) - - compress_context(ctx, strategy="dedup_then_truncate", max_messages=10) - - messages = ctx.stm_get() - assert len(messages) == 1 - assert messages[0].data == {"values": {1, 2, 3}} - - -def test_compress_context_target_tokens_trims_long_history() -> None: - ctx = Context(config=Config()) - for idx in range(8): - ctx.stm_add(Message(role="user", text=f"long-message-{idx}-" + "x" * 120)) - - before_count = len(ctx.stm_get()) - compress_context(ctx, strategy="truncate", max_messages=8, target_tokens=80) - after_messages = ctx.stm_get() - - assert len(after_messages) < before_count - assert len(after_messages) >= 1 - - -def test_compress_context_negative_max_messages_keeps_unbounded_semantics() -> None: - ctx = Context(config=Config()) - for idx in range(4): - ctx.stm_add(Message(role="user", text=f"msg-{idx}")) - - before_messages = list(ctx.stm_get()) - compress_context( - ctx, - strategy="truncate", - max_messages=-1, - target_tokens=10_000, - ) - after_messages = ctx.stm_get() - - assert [message.text for message in after_messages] == [ - message.text for message in before_messages - ] - - -def test_compress_context_annotate_preserves_message_identity_fields() -> None: - ctx = Context(config=Config()) - ctx.stm_add( - Message( - role="assistant", - text="keep identity", - id="assistant-1", - mark=MessageMark.PERSISTENT, - ) - ) - ctx.stm_add(Message(role="user", text="latest")) - - compress_context(ctx, strategy="dedup_then_truncate", max_messages=1, phase="pre_tool") - - head = ctx.stm_get()[0] - assert head.id == "assistant-1" - assert head.mark == MessageMark.PERSISTENT - assert head.metadata.get("compressed") is True - assert head.metadata.get("strategy") == "dedup_then_truncate" - - -def test_compress_context_annotate_preserves_structured_message_fields() -> None: - ctx = Context(config=Config()) - ctx.stm_add( - Message( - role="assistant", - kind=MessageKind.CHAT, - text="keep attachment", - attachments=[AttachmentRef(kind=AttachmentKind.IMAGE, uri="https://example.com/a.png")], - data={"tool_calls": [{"id": "tc_1"}]}, - metadata={"trace": "1"}, - mark=MessageMark.PERSISTENT, - ) - ) - ctx.stm_add(Message(role="user", text="latest")) - - compress_context(ctx, strategy="dedup_then_truncate", max_messages=1, phase="pre_tool") - - head = ctx.stm_get()[0] - assert head.text == "keep attachment" - assert len(head.attachments) == 1 - assert head.attachments[0].uri == "https://example.com/a.png" - assert head.data == {"tool_calls": [{"id": "tc_1"}]} - assert head.metadata.get("trace") == "1" - assert head.metadata.get("compressed") is True - - -def test_compress_context_max_messages_preserves_protected_marks() -> None: - ctx = Context(config=Config()) - ctx.stm_add( - Message( - role="system", - text="immutable", - id="imm-1", - mark=MessageMark.IMMUTABLE, - ) - ) - ctx.stm_add( - Message( - role="assistant", - text="persistent", - id="persist-1", - mark=MessageMark.PERSISTENT, - ) - ) - for idx in range(4): - ctx.stm_add(Message(role="user", text=f"temp-{idx}", id=f"tmp-{idx}")) - - compress_context(ctx, strategy="truncate", max_messages=3, target_tokens=10_000) - - messages = ctx.stm_get() - ids = [message.id for message in messages] - assert "imm-1" in ids - assert "persist-1" in ids - assert len(messages) == 3 diff --git a/tests/unit/test_context_implementation.py b/tests/unit/test_context_implementation.py index 43488ef4..a7d20770 100644 --- a/tests/unit/test_context_implementation.py +++ b/tests/unit/test_context_implementation.py @@ -1,5 +1,7 @@ +from __future__ import annotations import pytest + from dare_framework.config import Config from dare_framework.context.context import Context from dare_framework.context.types import AttachmentKind, AttachmentRef, Budget, Message @@ -7,9 +9,11 @@ from dare_framework.tool._internal.tools.noop_tool import NoopTool from dare_framework.tool.tool_manager import ToolManager -def test_context_initialization(): + +def test_context_initialization() -> None: config = Config() ctx = Context(id="test-id", config=config) + assert ctx.id == "test-id" assert isinstance(ctx.budget, Budget) assert ctx.short_term_memory is not None @@ -18,31 +22,35 @@ def test_context_initialization(): assert ctx.config is config assert ctx.sys_prompt is None -def test_context_stm_methods(): + +def test_context_stm_methods() -> None: ctx = Context(config=Config()) msg = Message(role="user", text="hello") ctx.stm_add(msg) - + messages = ctx.stm_get() assert len(messages) == 1 assert messages[0].text == "hello" - + ctx.stm_clear() assert len(ctx.stm_get()) == 0 -def test_context_budget_methods(): + +def test_context_budget_methods() -> None: ctx = Context(config=Config(), budget=Budget(max_tokens=100)) ctx.budget_use("tokens", 50) + assert ctx.budget.used_tokens == 50 assert ctx.budget_remaining("tokens") == 50 - - ctx.budget_check() # Should not raise - + + ctx.budget_check() + ctx.budget_use("tokens", 60) with pytest.raises(RuntimeError, match="Token budget exceeded"): ctx.budget_check() -def test_context_assemble(): + +def test_context_assemble() -> None: prompt = Prompt( prompt_id="test.system", role="system", @@ -52,8 +60,9 @@ def test_context_assemble(): ) ctx = Context(config=Config(), sys_prompt=prompt) ctx.stm_add(Message(role="user", text="hi")) - + assembled = ctx.assemble() + assert assembled.sys_prompt is not None assert assembled.sys_prompt.content == "You are a helpful assistant" assert len(assembled.messages) == 1 @@ -85,12 +94,12 @@ def test_context_assemble_preserves_chat_attachments() -> None: assert assembled.messages[0].attachments[0].uri == "https://example.com/a.png" -def test_context_requires_non_null_config(): +def test_context_requires_non_null_config() -> None: with pytest.raises(ValueError, match="non-null Config"): Context(id="missing-config", config=None) # type: ignore[arg-type] -def test_context_list_tools_returns_capability_descriptors_from_tool_manager(): +def test_context_list_tools_returns_capability_descriptors_from_tool_manager() -> None: manager = ToolManager(load_entrypoints=False) manager.register_tool(NoopTool()) ctx = Context(config=Config(), tool_gateway=manager) @@ -102,7 +111,7 @@ def test_context_list_tools_returns_capability_descriptors_from_tool_manager(): assert tools[0].name == "noop" -def test_context_exposes_public_tool_gateway_accessor_and_setter(): +def test_context_exposes_public_tool_gateway_accessor_and_setter() -> None: manager = ToolManager(load_entrypoints=False) ctx = Context(config=Config()) @@ -113,15 +122,12 @@ def test_context_exposes_public_tool_gateway_accessor_and_setter(): class _FakeRetrieval: - def __init__(self, messages: list[Message], *, fail: bool = False) -> None: + def __init__(self, messages: list[Message]) -> None: self._messages = list(messages) - self._fail = fail self.calls: list[tuple[str, dict[str, object]]] = [] def get(self, query: str = "", **kwargs: object) -> list[Message]: self.calls.append((query, dict(kwargs))) - if self._fail: - raise RuntimeError("retrieval failed") return list(self._messages) def add(self, message: Message) -> None: @@ -135,29 +141,7 @@ def compress(self, **kwargs: object) -> int: return 0 -class _CompressionRecordingSTM(_FakeRetrieval): - def __init__(self, messages: list[Message]) -> None: - super().__init__(messages) - self.compress_calls: list[dict[str, object]] = [] - - def compress(self, **kwargs: object) -> int: - self.compress_calls.append(dict(kwargs)) - raw_limit = kwargs.get("max_messages") - limit = raw_limit if isinstance(raw_limit, int) and raw_limit >= 0 else None - if limit is not None and len(self._messages) > limit: - self._messages = self._messages[-limit:] - return 0 - - -class _RecordingMovingCompressor: - def __init__(self) -> None: - self.calls: list[dict[str, object]] = [] - - async def prune(self, context: Context, **options: object) -> None: - self.calls.append({"context": context, "options": dict(options)}) - - -def test_context_assemble_fuses_ltm_and_knowledge_with_latest_user_query(): +def test_context_assemble_ignores_optional_retrieval_sources_in_default_strategy() -> None: ltm = _FakeRetrieval([Message(role="assistant", text="ltm-hit")]) knowledge = _FakeRetrieval([Message(role="assistant", text="knowledge-hit")]) ctx = Context( @@ -165,333 +149,57 @@ def test_context_assemble_fuses_ltm_and_knowledge_with_latest_user_query(): long_term_memory=ltm, knowledge=knowledge, ) - ctx.stm_add(Message(role="user", text="old request")) - ctx.stm_add(Message(role="assistant", text="ack")) ctx.stm_add(Message(role="user", text="latest request")) assembled = ctx.assemble() - contents = [message.text for message in assembled.messages] - assert contents == ["old request", "ack", "latest request", "ltm-hit", "knowledge-hit"] - assert ltm.calls and ltm.calls[0][0] == "latest request" - assert knowledge.calls and knowledge.calls[0][0] == "latest request" - assert assembled.metadata["retrieval"]["ltm_count"] == 1 - assert assembled.metadata["retrieval"]["knowledge_count"] == 1 - assert assembled.metadata["retrieval"]["degraded"] is False - - -def test_context_assemble_degrades_when_token_budget_low(): - ltm = _FakeRetrieval([Message(role="assistant", text="ltm-hit")]) - knowledge = _FakeRetrieval([Message(role="assistant", text="knowledge-hit")]) - # Force a low remaining token budget so retrieval should be skipped. - budget = Budget(max_tokens=32) - ctx = Context( - config=Config(), - budget=budget, - long_term_memory=ltm, - knowledge=knowledge, - ) - ctx.stm_add(Message(role="user", text="x" * 160)) - - assembled = ctx.assemble() - - contents = [message.text for message in assembled.messages] - assert contents == ["x" * 160] - assert assembled.metadata["retrieval"]["degraded"] is True - assert assembled.metadata["retrieval"]["degrade_reason"] == "token_budget_low" - - -def test_context_assemble_zero_token_budget_skips_retrieval() -> None: - ltm = _FakeRetrieval([Message(role="assistant", text="ltm-hit")]) - knowledge = _FakeRetrieval([Message(role="assistant", text="knowledge-hit")]) - ctx = Context( - config=Config(), - budget=Budget(max_tokens=0), - long_term_memory=ltm, - knowledge=knowledge, - ) - ctx.stm_add(Message(role="user", text="query")) - - assembled = ctx.assemble() - - assert [message.text for message in assembled.messages] == ["query"] - assert assembled.metadata["retrieval"]["degraded"] is True - assert assembled.metadata["retrieval"]["degrade_reason"] == "token_budget_low" + assert [message.text for message in assembled.messages] == ["latest request"] assert ltm.calls == [] assert knowledge.calls == [] + assert assembled.metadata == {"context_id": ctx.id} -def test_context_assemble_handles_retrieval_exception_gracefully(): - ltm = _FakeRetrieval([Message(role="assistant", text="ltm-hit")], fail=True) - knowledge = _FakeRetrieval([Message(role="assistant", text="knowledge-hit")]) - ctx = Context( - config=Config(), - long_term_memory=ltm, - knowledge=knowledge, - ) - ctx.stm_add(Message(role="user", text="query")) - - assembled = ctx.assemble() - - contents = [message.text for message in assembled.messages] - assert contents == ["query", "knowledge-hit"] - assert assembled.metadata["retrieval"]["degraded"] is True - assert assembled.metadata["retrieval"]["degrade_reason"] == "ltm_retrieval_failed" - - -def test_context_assemble_single_source_uses_full_retrieval_budget(): - ltm = _FakeRetrieval([Message(role="assistant", text="x" * 64)]) - config = Config( - long_term_memory={ - "assemble_top_k": 1, - "assemble_reserve_tokens": 0, - "assemble_ratio": 0.5, - }, - knowledge={ - "assemble_top_k": 1, - "assemble_ratio": 0.5, - }, - ) - # Remaining retrieval budget ~= 40 tokens after STM estimate. - ctx = Context( - config=config, - budget=Budget(max_tokens=49), - long_term_memory=ltm, - knowledge=None, - ) - ctx.stm_add(Message(role="user", text="q")) - - assembled = ctx.assemble() - - contents = [message.text for message in assembled.messages] - assert contents == ["q", "x" * 64] - assert assembled.metadata["retrieval"]["ltm_count"] == 1 - assert assembled.metadata["retrieval"]["degraded"] is False - - -def test_context_assemble_skips_oversized_retrieval_hits_and_keeps_later_candidates(): - ltm = _FakeRetrieval( - [ - Message(role="assistant", text="x" * 220), - Message(role="assistant", text="small-hit"), - ] - ) - config = Config( - long_term_memory={ - "assemble_top_k": 2, - "assemble_reserve_tokens": 0, - "assemble_ratio": 1.0, - }, - knowledge={ - "assemble_top_k": 0, - "assemble_ratio": 0.0, - }, - ) - ctx = Context( - config=config, - budget=Budget(max_tokens=35), - long_term_memory=ltm, - knowledge=None, - ) - ctx.stm_add(Message(role="user", text="q")) - - assembled = ctx.assemble() - - contents = [message.text for message in assembled.messages] - assert contents == ["q", "small-hit"] - assert assembled.metadata["retrieval"]["ltm_count"] == 1 - assert assembled.metadata["retrieval"]["degraded"] is True - - -def test_context_assemble_reserve_tokens_respects_knowledge_only_config(): - knowledge = _FakeRetrieval([Message(role="assistant", text="x" * 64)]) - config = Config( - long_term_memory={"assemble_top_k": 0}, - knowledge={ - "assemble_top_k": 1, - "assemble_ratio": 1.0, - "assemble_reserve_tokens": 0, - }, - ) - ctx = Context( - config=config, - budget=Budget(max_tokens=40), - long_term_memory=None, - knowledge=knowledge, - ) - ctx.stm_add(Message(role="user", text="q")) - - assembled = ctx.assemble() - - contents = [message.text for message in assembled.messages] - assert contents == ["q", "x" * 64] - assert assembled.metadata["retrieval"]["knowledge_count"] == 1 - assert assembled.metadata["retrieval"]["degraded"] is False - - -def test_context_assemble_ignores_inactive_ltm_reserve_tokens_for_knowledge_only_retrieval(): - knowledge = _FakeRetrieval([Message(role="assistant", text="x" * 64)]) - config = Config( - long_term_memory={ - "assemble_top_k": 0, - "assemble_reserve_tokens": 10_000, - }, - knowledge={ - "assemble_top_k": 1, - "assemble_ratio": 1.0, - "assemble_reserve_tokens": 0, - }, - ) - ctx = Context( - config=config, - budget=Budget(max_tokens=40), - long_term_memory=None, - knowledge=knowledge, - ) - ctx.stm_add(Message(role="user", text="q")) - - assembled = ctx.assemble() - - contents = [message.text for message in assembled.messages] - assert contents == ["q", "x" * 64] - assert assembled.metadata["retrieval"]["knowledge_count"] == 1 - assert assembled.metadata["retrieval"]["degraded"] is False - - -def test_context_assemble_rebalances_budget_when_ltm_retrieval_fails(): - ltm = _FakeRetrieval([Message(role="assistant", text="ltm-hit")], fail=True) - knowledge = _FakeRetrieval([Message(role="assistant", text="x" * 64)]) - config = Config( - long_term_memory={ - "assemble_top_k": 1, - "assemble_ratio": 0.5, - "assemble_reserve_tokens": 0, - }, - knowledge={ - "assemble_top_k": 1, - "assemble_ratio": 0.5, - "assemble_reserve_tokens": 0, - }, - ) - ctx = Context( - config=config, - budget=Budget(max_tokens=34), - long_term_memory=ltm, - knowledge=knowledge, - ) - ctx.stm_add(Message(role="user", text="q")) - - assembled = ctx.assemble() - - contents = [message.text for message in assembled.messages] - assert contents == ["q", "x" * 64] - assert assembled.metadata["retrieval"]["ltm_count"] == 0 - assert assembled.metadata["retrieval"]["knowledge_count"] == 1 - assert assembled.metadata["retrieval"]["degraded"] is True - assert assembled.metadata["retrieval"]["degrade_reason"] == "ltm_retrieval_failed" - - -def test_context_assemble_skips_zero_budget_source_retrieval_call() -> None: - ltm = _FakeRetrieval([Message(role="assistant", text="ltm-hit")]) - knowledge = _FakeRetrieval([Message(role="assistant", text="knowledge-hit")]) - config = Config( - long_term_memory={ - "assemble_top_k": 1, - "assemble_ratio": 0.0, - "assemble_reserve_tokens": 0, - }, - knowledge={ - "assemble_top_k": 1, - "assemble_ratio": 1.0, - "assemble_reserve_tokens": 0, - }, - ) - ctx = Context( - config=config, - budget=Budget(max_tokens=60), - long_term_memory=ltm, - knowledge=knowledge, - ) - ctx.stm_add(Message(role="user", text="q")) - - assembled = ctx.assemble() +class _RecordingMovingCompressor: + def __init__(self) -> None: + self.calls: list[dict[str, object]] = [] - contents = [message.text for message in assembled.messages] - assert contents == ["q", "knowledge-hit"] - assert ltm.calls == [] - assert len(knowledge.calls) == 1 - assert assembled.metadata["retrieval"]["degraded"] is False + async def prune(self, context: Context, **options: object) -> None: + self.calls.append({"context": context, "options": dict(options)}) -def test_context_assemble_handles_overflowing_numeric_retrieval_config() -> None: - ltm = _FakeRetrieval([Message(role="assistant", text="ltm-hit")]) - config = Config( - long_term_memory={ - "assemble_top_k": float("inf"), - "assemble_reserve_tokens": float("inf"), - }, - knowledge={"assemble_top_k": 0}, - ) - ctx = Context( - config=config, - long_term_memory=ltm, - knowledge=None, - ) - ctx.stm_add(Message(role="user", text="query")) +@pytest.mark.asyncio +async def test_context_compress_without_moving_compressor_is_noop() -> None: + ctx = Context(config=Config()) + ctx.stm_add(Message(role="user", text="hello")) - assembled = ctx.assemble() + await ctx.compress(max_context_tokens=128) - assert assembled.metadata["retrieval"]["ltm_requested"] == 3 + assert [message.text for message in ctx.stm_get()] == ["hello"] -def test_context_assemble_rejects_infinite_ratio_and_keeps_budget_guardrails() -> None: - ltm = _FakeRetrieval([Message(role="assistant", text="x" * 220)]) - config = Config( - long_term_memory={ - "assemble_top_k": 1, - "assemble_ratio": float("inf"), - "assemble_reserve_tokens": 0, - }, - knowledge={"assemble_top_k": 0}, - ) - ctx = Context( - config=config, - budget=Budget(max_tokens=40), - long_term_memory=ltm, - knowledge=None, - ) - ctx.stm_add(Message(role="user", text="q")) +@pytest.mark.asyncio +async def test_context_compress_uses_context_window_tokens_when_present() -> None: + ctx = Context(config=Config(), context_window_tokens=256) + compressor = _RecordingMovingCompressor() + ctx.set_moving_compressor(compressor) - assembled = ctx.assemble() + await ctx.compress() - contents = [message.text for message in assembled.messages] - assert contents == ["q"] - assert assembled.metadata["retrieval"]["ltm_count"] == 0 - assert assembled.metadata["retrieval"]["degraded"] is True + assert len(compressor.calls) == 1 + assert compressor.calls[0]["context"] is ctx + assert compressor.calls[0]["options"] == {"max_context_tokens": 256} -def test_context_assemble_handles_overflowing_numeric_ratio_config() -> None: - ltm = _FakeRetrieval([Message(role="assistant", text="ltm-hit")]) - config = Config( - long_term_memory={ - "assemble_top_k": 1, - "assemble_ratio": 10**10000, - "assemble_reserve_tokens": 0, - }, - knowledge={"assemble_top_k": 0}, - ) - ctx = Context( - config=config, - long_term_memory=ltm, - knowledge=None, - ) - ctx.stm_add(Message(role="user", text="query")) +@pytest.mark.asyncio +async def test_context_compress_prefers_explicit_max_context_tokens() -> None: + ctx = Context(config=Config(), context_window_tokens=256) + compressor = _RecordingMovingCompressor() + ctx.set_moving_compressor(compressor) - assembled = ctx.assemble() + await ctx.compress(max_context_tokens=64) - contents = [message.text for message in assembled.messages] - assert contents == ["query", "ltm-hit"] - assert assembled.metadata["retrieval"]["ltm_count"] == 1 + assert len(compressor.calls) == 1 + assert compressor.calls[0]["options"] == {"max_context_tokens": 64} @pytest.mark.asyncio @@ -507,69 +215,3 @@ async def test_context_assemble_for_model_runs_moving_compressor_with_context_wi assert compressor.calls[0]["context"] is ctx assert compressor.calls[0]["options"] == {"max_context_tokens": 256} assert [message.text for message in assembled.messages] == ["query"] - - -def test_context_compress_max_messages_uses_backend_compress_only(monkeypatch: pytest.MonkeyPatch) -> None: - stm = _CompressionRecordingSTM( - [ - Message(role="user", text="m0"), - Message(role="assistant", text="m1"), - Message(role="user", text="m2"), - ] - ) - ctx = Context(config=Config(), short_term_memory=stm) - - def _unexpected_compress_context(*args: object, **kwargs: object) -> None: - _ = (args, kwargs) - raise AssertionError("compress_context should not be called for basic max_messages compression") - - monkeypatch.setattr( - "dare_framework.compression.core.compress_context", - _unexpected_compress_context, - ) - - ctx.compress(max_messages=2) - - assert len(stm.compress_calls) == 1 - assert stm.compress_calls[0].get("max_messages") == 2 - assert [message.text for message in ctx.stm_get()] == ["m1", "m2"] - - -def test_context_compress_zero_max_messages_clears_default_stm() -> None: - ctx = Context(config=Config()) - ctx.stm_add(Message(role="user", text="m0")) - ctx.stm_add(Message(role="assistant", text="m1")) - - ctx.compress(max_messages=0) - - assert ctx.stm_get() == [] - - -def test_context_compress_advanced_path_preserves_backend_semantics(monkeypatch: pytest.MonkeyPatch) -> None: - stm = _CompressionRecordingSTM( - [ - Message(role="user", text="m0"), - Message(role="assistant", text="m1"), - Message(role="user", text="m2"), - ] - ) - ctx = Context(config=Config(), short_term_memory=stm) - calls: list[dict[str, object]] = [] - stm_sizes_seen_by_strategy: list[int] = [] - - def _record_compress_context(context: Context, **kwargs: object) -> None: - stm_sizes_seen_by_strategy.append(len(context.stm_get())) - calls.append(dict(kwargs)) - - monkeypatch.setattr( - "dare_framework.compression.core.compress_context", - _record_compress_context, - ) - - ctx.compress(max_messages=2, target_tokens=100, strategy="truncate", tool_pair_safe=True) - - assert len(stm.compress_calls) == 1 - assert stm.compress_calls[0].get("max_messages") == 2 - assert calls and calls[0].get("max_messages") == 2 - assert stm_sizes_seen_by_strategy == [3] - assert [message.text for message in ctx.stm_get()] == ["m1", "m2"] diff --git a/tests/unit/test_example_10_agentscope_compat.py b/tests/unit/test_example_10_agentscope_compat.py index 974dbf7c..a5ba4e61 100644 --- a/tests/unit/test_example_10_agentscope_compat.py +++ b/tests/unit/test_example_10_agentscope_compat.py @@ -263,6 +263,8 @@ def test_truncated_formatter_truncates_and_preserves_tool_pairs() -> None: assert tool_result_names.issubset(tool_call_ids) + + def test_json_session_bridge_roundtrip(tmp_path: Path) -> None: module = _load_example_module() notebook = module.CompatPlanNotebook() diff --git a/tests/unit/test_react_agent_gateway_injection.py b/tests/unit/test_react_agent_gateway_injection.py index 07ebc2f1..1304efa2 100644 --- a/tests/unit/test_react_agent_gateway_injection.py +++ b/tests/unit/test_react_agent_gateway_injection.py @@ -6,10 +6,8 @@ from dare_framework.agent.react_agent import ReactAgent from dare_framework.config import Config -from dare_framework.context import Context -from dare_framework.context.manage_context import MANAGE_CONTEXT_TOOL_NAME +from dare_framework.context import Context, Message from dare_framework.context.types import MessageKind -from dare_framework.context.smartcontext import SmartContext from dare_framework.model.types import ModelInput, ModelResponse from dare_framework.tool.types import CapabilityDescriptor, CapabilityType, ToolResult @@ -153,42 +151,6 @@ async def generate(self, model_input: ModelInput, *, options: Any | None = None) return response -class _CompressionRecordingContext(Context): - def __init__(self, *, config: Config) -> None: - super().__init__(config=config) - self.compress_calls: list[dict[str, Any]] = [] - - def compress(self, **options: Any) -> None: - self.compress_calls.append(dict(options)) - super().compress(**options) - - -class _CompressionRecordingSmartContext(SmartContext): - def __init__(self, *, config: Config) -> None: - super().__init__(config=config) - self.compress_calls: list[dict[str, Any]] = [] - - def compress(self, **options: Any) -> None: - self.compress_calls.append(dict(options)) - super().compress(**options) - - -class _FinalOnlyModel: - async def generate(self, model_input: ModelInput, *, options: Any | None = None) -> ModelResponse: - _ = (model_input, options) - return ModelResponse(content="final", tool_calls=[]) - - -class _CapturingFinalModel: - def __init__(self) -> None: - self.last_messages: list[Any] | None = None - - async def generate(self, model_input: ModelInput, *, options: Any | None = None) -> ModelResponse: - _ = options - self.last_messages = list(model_input.messages) - return ModelResponse(content="final", tool_calls=[]) - - class _NonConvergingToolModel: def __init__(self) -> None: self._idx = 0 @@ -208,21 +170,6 @@ async def generate(self, model_input: ModelInput, *, options: Any | None = None) ) -class _ManageContextGateway(_RecordingGateway): - def __init__(self) -> None: - super().__init__("manage-context") - self._capabilities = [ - CapabilityDescriptor( - id=MANAGE_CONTEXT_TOOL_NAME, - type=CapabilityType.TOOL, - name=MANAGE_CONTEXT_TOOL_NAME, - description="manage context", - input_schema={"type": "object"}, - output_schema={"type": "object"}, - ) - ] - - @pytest.mark.asyncio async def test_react_agent_prefers_injected_gateway_over_context_gateway() -> None: context_gateway = _RecordingGateway("context") @@ -300,7 +247,7 @@ async def test_react_agent_emits_intermediate_transport_events_in_order() -> Non tool_gateway=gateway, ) - result = await agent.execute("test", transport=transport) + result = await agent.execute(Message(role="user", text="test"), transport=transport) assert result.success is True message_kinds = [envelope.payload.message_kind for envelope in transport.sent] @@ -332,7 +279,7 @@ async def test_react_agent_transport_loop_emits_single_terminal_result_event() - ) await agent._execute_polled_message( - "test", + Message(role="user", text="test"), channel=transport, envelope_id="req_1", ) @@ -363,7 +310,7 @@ async def test_react_agent_emits_terminal_message_for_repeated_tool_guard() -> N tool_gateway=gateway, ) - result = await agent.execute("test", transport=transport) + result = await agent.execute(Message(role="user", text="test"), transport=transport) assert result.success is True assert transport.sent @@ -386,7 +333,7 @@ async def test_react_agent_emits_terminal_message_for_max_round_exit() -> None: max_tool_rounds=2, ) - result = await agent.execute("test", transport=transport) + result = await agent.execute(Message(role="user", text="test"), transport=transport) assert result.success is True assert transport.sent @@ -395,120 +342,6 @@ async def test_react_agent_emits_terminal_message_for_max_round_exit() -> None: assert "达到最大轮次" in str(last_envelope.payload.data["output"]) -@pytest.mark.asyncio -async def test_react_agent_auto_compress_triggers_before_model_call() -> None: - context = _CompressionRecordingContext(config=Config()) - context.budget.max_tokens = 100 - gateway = _RecordingGateway("injected") - agent = ReactAgent( - name="react-test-auto-compress", - model=_FinalOnlyModel(), - context=context, - tool_gateway=gateway, - auto_compress=True, - compress_trigger_ratio=0.01, - compress_target_ratio=0.5, - ) - - result = await agent("test auto compress trigger") - - assert result.success is True - assert len(context.compress_calls) >= 1 - first_call = context.compress_calls[0] - assert first_call.get("tool_pair_safe") is True - assert first_call.get("target_tokens") is not None - - -@pytest.mark.asyncio -async def test_react_agent_auto_compress_nan_ratios_fallback_to_defaults() -> None: - context = _CompressionRecordingContext(config=Config()) - context.budget.max_tokens = 100 - gateway = _RecordingGateway("injected") - agent = ReactAgent( - name="react-test-auto-compress-nan-ratios", - model=_FinalOnlyModel(), - context=context, - tool_gateway=gateway, - auto_compress=True, - compress_trigger_ratio=float("nan"), - compress_target_ratio=float("nan"), - ) - - result = await agent("x" * 600) - - assert result.success is True - assert len(context.compress_calls) >= 1 - first_call = context.compress_calls[0] - assert first_call.get("target_tokens") == 75 - - -@pytest.mark.asyncio -async def test_react_agent_without_auto_compress_keeps_legacy_behavior() -> None: - context = _CompressionRecordingContext(config=Config()) - gateway = _RecordingGateway("injected") - agent = ReactAgent( - name="react-test-no-auto-compress", - model=_FinalOnlyModel(), - context=context, - tool_gateway=gateway, - auto_compress=False, - ) - - result = await agent("test no auto compress") - - assert result.success is True - assert context.compress_calls == [] - - -@pytest.mark.asyncio -async def test_react_agent_auto_compress_triggers_in_smart_context_path() -> None: - context = _CompressionRecordingSmartContext(config=Config()) - context.budget.max_tokens = 100 - gateway = _RecordingGateway("injected") - agent = ReactAgent( - name="react-test-smartcontext-auto-compress", - model=_FinalOnlyModel(), - context=context, - tool_gateway=gateway, - auto_compress=True, - compress_trigger_ratio=0.01, - compress_target_ratio=0.5, - ) - - result = await agent("smart context compress") - - assert result.success is True - assert len(context.compress_calls) >= 1 - assert context.compress_calls[0].get("tool_pair_safe") is True - - -@pytest.mark.asyncio -async def test_react_agent_auto_compress_reappends_reflection_prompt_in_smart_context_path() -> None: - context = _CompressionRecordingSmartContext(config=Config()) - context.budget.max_tokens = 100 - context.update_task_complete(True) - gateway = _ManageContextGateway() - model = _CapturingFinalModel() - agent = ReactAgent( - name="react-test-smartcontext-reflection-prompt", - model=model, - context=context, - tool_gateway=gateway, - auto_compress=True, - compress_trigger_ratio=0.01, - compress_target_ratio=0.5, - ) - - result = await agent("smart context compress") - - assert result.success is True - assert model.last_messages is not None - assert any( - (message.text or "") == "【提示】请先调用 manage_context 根据任务初始化 context 状态。" - for message in model.last_messages - ) - - @pytest.mark.asyncio async def test_react_agent_loop_guard_emits_terminal_message_event() -> None: context = Context(config=Config()) @@ -522,7 +355,7 @@ async def test_react_agent_loop_guard_emits_terminal_message_event() -> None: max_tool_rounds=10, ) - result = await agent.execute("test loop guard", transport=transport) + result = await agent.execute(Message(role="user", text="test loop guard"), transport=transport) assert result.success is True assert transport.sent[-1].payload.message_kind is MessageKind.CHAT @@ -542,7 +375,7 @@ async def test_react_agent_max_round_exit_emits_terminal_message_event() -> None max_tool_rounds=2, ) - result = await agent.execute("test max rounds", transport=transport) + result = await agent.execute(Message(role="user", text="test max rounds"), transport=transport) assert result.success is True assert transport.sent[-1].payload.message_kind is MessageKind.CHAT